From 40362e46e2d04a93377cf542e6b5a303130dce88 Mon Sep 17 00:00:00 2001 From: bin913 <842884726@qq.com> Date: Tue, 16 Jun 2026 10:48:34 +0800 Subject: [PATCH 1/6] add run_op.sh --- tools/run_op.sh | 132 ++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 132 insertions(+) create mode 100755 tools/run_op.sh diff --git a/tools/run_op.sh b/tools/run_op.sh new file mode 100755 index 00000000..b0e219a5 --- /dev/null +++ b/tools/run_op.sh @@ -0,0 +1,132 @@ +#!/bin/bash + +# FlagBLAS operator test runner script +# Reference: FlagGems/tools/test-op.sh +# Usage: ./tools/run_op.sh [PR_ID] +# Environment variables: +# CHANGED_FILES - list of changed files (space-separated), or "__ALL__" for full test + +PR_ID=$1 + +# Leave this for debugging's purpose +echo "PR_ID=${PR_ID}" + +COLLECT_COVERAGE="" +FAIL_FAST=false + +if [[ "$CHANGED_FILES" == "__ALL__" ]]; then + # Replace "__ALL__" with all tests + CHANGED_FILES=$(find tests -name "test*.py") + # add options to generate summary report + EXTRA_OPTS="--md-report" + EXTRA_OPTS+=" --md-report-verbose=1" + EXTRA_OPTS+=" --md-report-output=${PR_ID}-summary.md" + SUFFIX="" + COLLECT_COVERAGE="yes" +else + # for per-PR test, fail early + FAIL_FAST=true + EXTRA_OPTS="-x" + SUFFIX="-${GITHUB_SHA::7}" +fi + +# Test cases that needs to run quick cpu tests +NO_QUICK_CPU_TESTS=( + "tests/conftest.py" + "tests/accuracy_utils.py" + "tests/__init__.py" +) + +# Extract test cases from CHANGED_FILES +TEST_CASES=() +PERF_TEST_CASES=() +TEST_CASES_CPU=() +for item in $CHANGED_FILES; do + file_name=$(basename "$item") + case $item in + tests/*.py) + if [[ "$file_name" == test*.py ]]; then + TEST_CASES+=($item) + fi + ;; + benchmark/test*) + PERF_TEST_CASES+=($item) + ;; + esac + + # filter out tests that do not need quick CPU mode tests + found=0 + for item_cpu in "${NO_QUICK_CPU_TESTS[@]}"; do + if [[ "$item" == "$item_cpu" ]]; then + found=1 + break + fi + done + if (( $found == 0 )); then + case $item in + tests/*.py) + if [[ "$file_name" == test*.py ]]; then + TEST_CASES_CPU+=($item) + fi + ;; + esac + fi +done + +# Skip tests if no tests file is found +if [[ ${#TEST_CASES[@]} -eq 0 && ${#PERF_TEST_CASES[@]} -eq 0 ]]; then + exit 0 +fi + +# Clear existing coverage data if any +coverage erase + +FAILURES=() +for item in "${TEST_CASES[@]}"; do + echo "Running unit tests for ${item}" + if ! coverage run -m pytest -s ${EXTRA_OPTS} ${item}; then + if $FAIL_FAST; then exit 1; fi + FAILURES+=("${item}") + fi +done + +# Run quick-cpu test if necessary +for item in "${TEST_CASES_CPU[@]}"; do + echo "Running quick-cpu mode unit tests for ${item}" + if ! coverage run -m pytest -s ${EXTRA_OPTS} ${item} --ref=cpu --quick; then + if $FAIL_FAST; then exit 1; fi + FAILURES+=("${item} (quick-cpu)") + fi +done + +# Run benchmark test if necessary +for item in "${PERF_TEST_CASES[@]}"; do + echo "Running benchmark tests for ${item}" + echo "pytest -s ${item} --level core --record log" + if ! pytest -s ${item} --level core --record log; then + if $FAIL_FAST; then exit 1; fi + FAILURES+=("${item} (benchmark)") + fi +done + +# Process coverage data only when full-range testing +# Coverage data HTML dumped to `htmlcov/` by default +if [ -n "$COLLECT_COVERAGE" ]; then + coverage combine + coverage html + rm -fr coverage + mkdir coverage + mv htmlcov coverage/ + echo "${PR_ID}${SUFFIX::7}" > coverage/COVERAGE_ID + mv ${PR_ID}-summary.md coverage/ut-summary.md +fi + +# Report failures +if [[ ${#FAILURES[@]} -gt 0 ]]; then + echo "" + echo "=== FAILED TESTS (${#FAILURES[@]}) ===" + for f in "${FAILURES[@]}"; do + echo " - ${f}" + done + exit 1 +fi \ No newline at end of file From ac0fe5343fd57589de731a597fb0d8dafc7408d5 Mon Sep 17 00:00:00 2001 From: bin913 <842884726@qq.com> Date: Tue, 16 Jun 2026 11:14:21 +0800 Subject: [PATCH 2/6] add run_op.sh in backend-test.yaml --- .github/workflows/backend-test.yaml | 19 +++++++++++++++---- 1 file changed, 15 insertions(+), 4 deletions(-) diff --git a/.github/workflows/backend-test.yaml b/.github/workflows/backend-test.yaml index 0c58f150..addd036a 100644 --- a/.github/workflows/backend-test.yaml +++ b/.github/workflows/backend-test.yaml @@ -17,9 +17,10 @@ on: type: string default: '' test_script: - description: 'Path to test script under tools/' - required: true + description: "Temporary parameter for alpha-ops test" + required: false type: string + default: '' changed_files: description: 'List of changed files in a PR' required: false @@ -75,9 +76,19 @@ jobs: run: | bash ${{ inputs.gpu_check_script }} - - name: Run backend tests + - name: Test alpha ops + if: ${{ inputs.test_script != '' }} env: CHANGED_FILES: ${{ inputs.changed_files }} shell: bash run: | - bash ${{ inputs.test_script }} ${{ inputs.vendor }} + bash ${{ inputs.test_script }} ${{ inputs.vendor }} ${{ inputs.pr_id }} + - name: Run backend tests + if: ${{ inputs.test_script == '' }} + shell: bash + env: + CHANGED_FILES: ${{ inputs.changed_files }} + run: | + source .venv/bin/activate + source tools/set-env.sh ${{ inputs.vendor }} + tools/run_op.sh ${{ inputs.pr_id }} From 74770c47460c7d9b8aa3a9f52a187e8ec30577cb Mon Sep 17 00:00:00 2001 From: bin913 <842884726@qq.com> Date: Thu, 18 Jun 2026 17:09:20 +0800 Subject: [PATCH 3/6] add static dispatch --- .../backend/_nvidia/hopper/ops/gemm.py | 106 ++++++++---------- src/flag_blas/runtime/dispatch.py | 54 +++++++++ 2 files changed, 98 insertions(+), 62 deletions(-) diff --git a/src/flag_blas/runtime/backend/_nvidia/hopper/ops/gemm.py b/src/flag_blas/runtime/backend/_nvidia/hopper/ops/gemm.py index b105da72..3033e291 100644 --- a/src/flag_blas/runtime/backend/_nvidia/hopper/ops/gemm.py +++ b/src/flag_blas/runtime/backend/_nvidia/hopper/ops/gemm.py @@ -31,7 +31,7 @@ _sgemm_tt_kernel, ) from flag_blas.runtime import torch_device_fn -from flag_blas.runtime.dispatch import SizeAutoDispatch +from flag_blas.runtime.dispatch import SizeAutoDispatch, StaticDispatch from flag_blas.utils import libentry, libtuner from flag_blas.utils.libentry import libcache @@ -1859,69 +1859,51 @@ def hgemm( ) aligned = _is_gemm_aligned(A, lda, B, ldb, C, ldc) - use_nn_kernel3 = aligned and (m * n > 2048 * 2048) and min(m, n) >= 64 with torch_device_fn.device(A.device): if transa == CUBLAS_OP_N and transb == CUBLAS_OP_N: - is_skinny = (m >= 16384 and max(n, k) <= 2048) or ( - n >= 16384 and max(m, k) <= 2048 - ) - if use_nn_kernel3 and is_skinny: - BLOCK_M = 128 - BLOCK_N = 256 - BLOCK_K = 64 - GROUP_M = 8 - NUM_STAGES = 4 - NUM_WARPS = 8 - NUM_CTAS = 1 - desc_a = TensorDescriptor( - base=A, - shape=[m, k], - strides=[lda, 1], - block_shape=[BLOCK_M, BLOCK_K], - ) - desc_b = TensorDescriptor( - base=B, - shape=[k, n], - strides=[ldb, 1], - block_shape=[BLOCK_K, BLOCK_N], - ) - desc_c = TensorDescriptor( - base=C, - shape=[m, n], - strides=[ldc, 1], - block_shape=[BLOCK_M, BLOCK_N], - ) - grid = (triton.cdiv(m, BLOCK_M) * triton.cdiv(n, BLOCK_N),) - _hgemm_nn_kernel4[grid]( - desc_a, - desc_b, - desc_c, - alpha, - beta, - m, - n, - k, - beta_is_zero, - BLOCK_M=BLOCK_M, - BLOCK_N=BLOCK_N, - BLOCK_K=BLOCK_K, - GROUP_M=GROUP_M, - num_stages=NUM_STAGES, - num_warps=NUM_WARPS, - num_ctas=NUM_CTAS, - ) - elif use_nn_kernel3: - _hgemm_nn_kernel3[grid]( - A, B, C, alpha, beta, m, n, k, lda, ldb, ldc, beta_is_zero - ) - elif aligned and max(m, n) <= 1024: - _hgemm_nn_kernel2[grid]( - A, B, C, alpha, beta, m, n, k, lda, ldb, ldc, beta_is_zero - ) - else: - _hgemm_nn_kernel[grid]( - A, B, C, alpha, beta, m, n, k, lda, ldb, ldc, beta_is_zero - ) + dispatch = StaticDispatch([ + # skinny + aligned large → kernel4 (TensorDescriptor, hardcoded config) + ( + lambda m, n, k, aligned, **_kw: + aligned and (m * n > 2048 * 2048) and min(m, n) >= 64 + and ((m >= 16384 and max(n, k) <= 2048) or (n >= 16384 and max(m, k) <= 2048)), + lambda: lambda: _hgemm_nn_kernel4[( + triton.cdiv(m, 128) * triton.cdiv(n, 256), + )]( + TensorDescriptor(base=A, shape=[m, k], strides=[lda, 1], block_shape=[128, 64]), + TensorDescriptor(base=B, shape=[k, n], strides=[ldb, 1], block_shape=[64, 256]), + TensorDescriptor(base=C, shape=[m, n], strides=[ldc, 1], block_shape=[128, 256]), + alpha, beta, m, n, k, beta_is_zero, + BLOCK_M=128, BLOCK_N=256, BLOCK_K=64, GROUP_M=8, + num_stages=4, num_warps=8, num_ctas=1, + ), + ), + # aligned large → kernel3 (TensorDescriptor, autotuned config) + ( + lambda m, n, k, aligned, **_kw: + aligned and (m * n > 2048 * 2048) and min(m, n) >= 64, + lambda: lambda: _hgemm_nn_kernel3[grid]( + A, B, C, alpha, beta, m, n, k, lda, ldb, ldc, beta_is_zero, + ), + ), + # aligned small → kernel2 (block_ptr) + ( + lambda m, n, k, aligned, **_kw: + aligned and max(m, n) <= 1024, + lambda: lambda: _hgemm_nn_kernel2[grid]( + A, B, C, alpha, beta, m, n, k, lda, ldb, ldc, beta_is_zero, + ), + ), + # default → kernel (original) + ( + lambda **_kw: True, + lambda: lambda: _hgemm_nn_kernel[grid]( + A, B, C, alpha, beta, m, n, k, lda, ldb, ldc, beta_is_zero, + ), + ), + ]) + runner = dispatch.lookup_and_build(m, n, k, aligned) + runner() elif transa == CUBLAS_OP_T and transb == CUBLAS_OP_N: is_skinny = (m >= 16384 and max(n, k) <= 2048) or ( n >= 16384 and max(m, k) <= 2048 diff --git a/src/flag_blas/runtime/dispatch.py b/src/flag_blas/runtime/dispatch.py index ea5d020d..1a1991f0 100644 --- a/src/flag_blas/runtime/dispatch.py +++ b/src/flag_blas/runtime/dispatch.py @@ -324,6 +324,60 @@ def _autotune( return runners.get(best_name) +class StaticDispatch: + """ + Developer-maintained static dispatch table. + + Maps shape conditions directly to kernel factories. No autotune, + no caching — the first matching entry wins immediately. + + Usage:: + + dispatch = StaticDispatch([ + (is_aligned, lambda: make_aligned_runner(A, lda, ...)), + (is_thin_large_m, lambda: make_thin_runner(A, lda, ...)), + (lambda **kw: True, lambda: make_fallback_runner(A, lda, ...)), + ]) + + runner = dispatch.lookup_and_build(m, n, k, aligned) + runner() + + Parameters + ---------- + table: + A list of ``(condition, factory)`` pairs evaluated **in order**. + Each ``condition`` is a callable with signature + ``(m, n, k, aligned, **extra) -> bool``. + Each ``factory`` is a zero-arg callable that returns a + ``Callable[[], None]`` runner. + The last entry should be a catch-all (condition always True). + """ + + _Entry = Tuple[Callable[..., bool], Callable[[], Callable[[], None]]] + + def __init__( + self, + table: List[_Entry], + ): + self._table = table + + def lookup_and_build( + self, + m: int, + n: int, + k: int, + aligned: bool, + **extra, + ) -> Callable[[], None]: + for condition, factory in self._table: + if condition(m=m, n=n, k=k, aligned=aligned, **extra): + return factory() + raise ValueError( + f"StaticDispatch: no matching entry for " + f"m={m}, n={n}, k={k}, aligned={aligned}" + ) + + class KernelRunner: """ A callable wrapper that executes a kernel with pre-bound arguments. From d0f6ae3bafd676bb6bf3af0be3f9234c019e6aa3 Mon Sep 17 00:00:00 2001 From: bin913 <842884726@qq.com> Date: Thu, 18 Jun 2026 19:50:38 +0800 Subject: [PATCH 4/6] add docs for static_dispatch.md --- docs/static_dispatch.md | 313 +++++++++++++++++++++++++++++++++++++ docs/static_dispatch_cn.md | 313 +++++++++++++++++++++++++++++++++++++ 2 files changed, 626 insertions(+) create mode 100644 docs/static_dispatch.md create mode 100644 docs/static_dispatch_cn.md diff --git a/docs/static_dispatch.md b/docs/static_dispatch.md new file mode 100644 index 00000000..7c1e10bd --- /dev/null +++ b/docs/static_dispatch.md @@ -0,0 +1,313 @@ +# StaticDispatch + +[Source: `src/flag_blas/runtime/dispatch.py`](../src/flag_blas/runtime/dispatch.py) + +`StaticDispatch` is a developer-maintained static dispatch table that maps shape conditions directly to pre-determined kernel factories. **No autotune, no benchmarking, no caching** — conditions are evaluated in order, and the first match wins immediately. + +--- + +## When to Use + +Use `StaticDispatch` when the **optimal shape → kernel mapping is already known** (e.g., through offline benchmarking, architectural analysis, or domain-specific heuristics). It provides: + +- **Predictable performance** — no runtime autotune overhead +- **Deterministic behavior** — same shape always maps to the same kernel +- **Minimal footprint** — no cache files, no DB I/O + +Contrast with `SizeAutoDispatch`, which benchmarks all candidates at runtime and persists the winner. + +--- + +## Core Concepts + +### Condition + +Signature: `(m: int, n: int, k: int, aligned: bool, **extra: Any) -> bool` + +A callable that determines whether the current shape matches this table entry. Conditions are evaluated **in table order**; the first that returns `True` wins. + +### Factory (Double-Lambda) + +Signature: `() -> Callable[[], None]` + +A zero-arg callable that returns a **runner**. The runner is itself a zero-arg callable that executes the kernel. + +**Critical**: When used with `@libentry()`-decorated Triton kernels, the factory must use a **double-lambda** wrapper: + +```python +# Correct: double lambda +lambda: lambda: kernel_fn[grid](arg1, arg2, ...) +# ^^^^ ^^^^^^^^^^^^^^^^^^^^^^^^ +# factory runner (deferred to runner()) + +# Wrong: single lambda — kernel executes immediately in factory() +lambda: kernel_fn[grid](arg1, arg2, ...) +``` + +**Why**: `kernel_fn[grid](args...)` calls `LibEntry.run()`, which launches the kernel immediately and returns a `(kernel, constexprs)` tuple — not a callable runner. The extra `lambda:` defers execution to `runner()` time. + +--- + +## Architecture + +Unlike `SizeAutoDispatch` with its multi-tier cache, `StaticDispatch` has a minimal design: + +``` +lookup_and_build(m, n, k, aligned, **extra) + │ + ├─ Entry 1: condition(m,n,k,aligned)? ─── True → factory() → runner + ├─ Entry 2: condition(m,n,k,aligned)? ─── True → factory() → runner + ├─ ... + └─ No match → raise ValueError +``` + +- **No filtering logic** — conditions encode all matching criteria inline (no separate `aligned`/`filter` params) +- **No cache** — every call re-evaluates conditions (cheap, just boolean logic) +- **Throw on miss** — if no entry matches, raises `ValueError` (the last entry should always be a catch-all) + +--- + +## API + +### Constructor + +```python +dispatch = StaticDispatch(table) +``` + +| Parameter | Type | Description | +|-----------|------|-------------| +| `table` | `List[Tuple[Condition, Factory]]` | Ordered list of `(condition, factory)` pairs. The **last entry must be a catch-all** (condition always returns `True`). | + +### lookup_and_build() + +```python +runner = dispatch.lookup_and_build(m, n, k, aligned, **extra) +runner() +``` + +| Parameter | Type | Description | +|-----------|------|-------------| +| `m` | `int` | M dimension | +| `n` | `int` | N dimension | +| `k` | `int` | K dimension | +| `aligned` | `bool` | Whether inputs are memory-aligned | +| `**extra` | — | Additional keyword arguments passed to each condition | + +**Returns**: `Callable[[], None]` — zero-arg runner; calling it executes the selected kernel. + +**Raises**: `ValueError` if no entry matches (should never happen with a proper catch-all). + +--- + +## Real-World Example: hgemm NN + +This example comes from [hopper/ops/gemm.py](../src/flag_blas/runtime/backend/_nvidia/hopper/ops/gemm.py), the `hgemm` function's NN branch: + +```python +from flag_blas.runtime.dispatch import StaticDispatch +from triton.tools.tensor_descriptor import TensorDescriptor + +dispatch = StaticDispatch([ + # ── Priority 1 (highest) ───────────────────────────────────── + # Skinny matrix with aligned large dimensions. + # Uses kernel4 with TensorDescriptor and hardcoded optimal config + # (no autotune needed — config proven optimal for this shape class). + ( + lambda m, n, k, aligned, **_kw: + aligned and (m * n > 2048 * 2048) and min(m, n) >= 64 + and ((m >= 16384 and max(n, k) <= 2048) + or (n >= 16384 and max(m, k) <= 2048)), + lambda: lambda: _hgemm_nn_kernel4[( + triton.cdiv(m, 128) * triton.cdiv(n, 256), + )]( + TensorDescriptor( + base=A, shape=[m, k], strides=[lda, 1], + block_shape=[128, 64], + ), + TensorDescriptor( + base=B, shape=[k, n], strides=[ldb, 1], + block_shape=[64, 256], + ), + TensorDescriptor( + base=C, shape=[m, n], strides=[ldc, 1], + block_shape=[128, 256], + ), + alpha, beta, m, n, k, beta_is_zero, + BLOCK_M=128, BLOCK_N=256, BLOCK_K=64, GROUP_M=8, + num_stages=4, num_warps=8, num_ctas=1, + ), + ), + + # ── Priority 2 ─────────────────────────────────────────────── + # Aligned + large dimensions (m×n > 4M, min dim ≥ 64). + # Uses kernel3 with TensorDescriptor and autotuned configs + # (its @libtuner picks the best BLOCK_M/N/K etc. at runtime). + ( + lambda m, n, k, aligned, **_kw: + aligned and (m * n > 2048 * 2048) and min(m, n) >= 64, + lambda: lambda: _hgemm_nn_kernel3[grid]( + A, B, C, alpha, beta, m, n, k, + lda, ldb, ldc, beta_is_zero, + ), + ), + + # ── Priority 3 ─────────────────────────────────────────────── + # Aligned + small/moderate dimensions (max ≤ 1024). + # Uses kernel2 with block_ptr and autotuned configs. + ( + lambda m, n, k, aligned, **_kw: + aligned and max(m, n) <= 1024, + lambda: lambda: _hgemm_nn_kernel2[grid]( + A, B, C, alpha, beta, m, n, k, + lda, ldb, ldc, beta_is_zero, + ), + ), + + # ── Priority 4 (catch-all) ─────────────────────────────────── + # Everything else: unaligned or moderate/large but not covered above. + # Uses the original level3 kernel with pointer-based access. + ( + lambda **_kw: True, + lambda: lambda: _hgemm_nn_kernel[grid]( + A, B, C, alpha, beta, m, n, k, + lda, ldb, ldc, beta_is_zero, + ), + ), +]) + +runner = dispatch.lookup_and_build(m, n, k, aligned) +runner() +``` + +### Dispatch Logic Summary + +| Priority | Condition | Kernel | +|----------|-----------|--------| +| 1 | aligned + large + skinny (one dim ≥ 16384, others ≤ 2048) | `kernel4` — TensorDescriptor, hardcoded config | +| 2 | aligned + large (m×n > 2048², min ≥ 64) | `kernel3` — TensorDescriptor, autotuned | +| 3 | aligned + small (max ≤ 1024) | `kernel2` — block_ptr, autotuned | +| 4 | everything else | `kernel` — original pointer-based, autotuned | + +### Kernel Variant Characteristics + +| Kernel | Source | Data Access | Config Strategy | Best For | +|--------|--------|-------------|-----------------|----------| +| `_hgemm_nn_kernel4` | Hopper (local) | `TensorDescriptor.load/store` + `int32` offsets | Hardcoded `(128,256,64,8,4,8,1)` | Skinny large matrices | +| `_hgemm_nn_kernel3` | Hopper (local) | `TensorDescriptor.load/store` | `@libtuner("hgemm_nn2")` | General large matrices | +| `_hgemm_nn_kernel2` | Hopper (local) | `tl.make_block_ptr` + `tl.advance` | `@libtuner("hgemm_nn")` | Small/moderate aligned | +| `_hgemm_nn_kernel` | Level3 (imported) | Raw pointers + `offs_{am,bn,k}` + mask logic | `@libtuner("hgemm_nn")` | Everything else | + +--- + +## Design Principles + +1. **No autotune**: The table is human-curated; no runtime benchmarking. +2. **Ordered matching**: Conditions evaluated top-to-bottom; first `True` wins. +3. **Catch-all required**: The last entry must match any shape (prevent `ValueError`). +4. **Mutually exclusive conditions**: Entries should not overlap to make behavior predictable. Higher-priority entries should have more specific conditions. +5. **Double-lambda factories**: Required when using `@libentry()`-decorated Triton kernels. The inner `lambda:` defers `LibEntry.run()` to `runner()` time. + +--- + +## Comparison with SizeAutoDispatch + +| | `SizeAutoDispatch` | `StaticDispatch` | +|---|---|---| +| **Selection** | Autotune (benchmark all → pick fastest) | Static conditions (first match wins) | +| **Cache** | In-memory + SQLite DB | None | +| **Persistence** | Cross-process via SQLite | None | +| **First call** | Benchmark cost (seconds) | Instant (condition evaluation only) | +| **Subsequent calls** | ~cache lookup | Instant (same as first call) | +| **Table building** | `add()` per variant | Single list in constructor | +| **Failure mode** | Fallback to first candidate | `ValueError` (mitigated by catch-all) | +| **Best for** | Unknown optimal mapping | Known optimal mapping | + +--- + +## Building a StaticDispatch Table + +### Step-by-step + +1. **Identify kernel variants**. List every kernel that could handle this operation, along with its strengths. + +2. **Write condition functions**. For each variant, define a lambda that returns `True` for the shapes where it shines. + +3. **Order by priority**. Put the most specific (narrowest condition) first, broadest last. + +4. **Add a catch-all**. The final entry must match everything (`lambda **_kw: True`). + +5. **Test edge cases**. Ensure shapes at condition boundaries route to the intended kernel. Use logging or debug prints during development. + +### Common Patterns + +**Priority by dimension**: +```python +[ + (lambda m, *_kw, **__: m > 8192, factory_a), # very large + (lambda m, *_kw, **__: m > 1024, factory_b), # large + (lambda m, *_kw, **__: m > 256, factory_c), # medium + (lambda **_kw: True, factory_d), # small +] +``` + +**Priority by alignment**: +```python +[ + (lambda aligned, **_kw: aligned and is_large(**kw), aligned_large_factory), + (lambda aligned, **_kw: aligned, aligned_small_factory), + (lambda **_kw: True, unaligned_factory), +] +``` + +**Combined dimensions + alignment** (as in hgemm_nn): +```python +[ + (lambda aligned, m, n, k, **_kw: + aligned and meets_criteria_A(m, n, k), + factory_a), + (lambda aligned, m, n, k, **_kw: + aligned and meets_criteria_B(m, n, k), + factory_b), + (lambda **_kw: True, + fallback_factory), +] +``` + +--- + +## Helper: KernelRunner + +```python +class KernelRunner: + def __init__(self, kernel: Callable, *args, **kwargs): ... + def __call__(self): + return self._kernel(*self._args, **self._kwargs) +``` + +A simple callable wrapper that binds a kernel function with its arguments. Useful when you don't need autotune, just a fixed implementation: + +```python +runner = KernelRunner(my_kernel_fn, A, B, C, alpha=1.0) +runner() # equivalent to my_kernel_fn(A, B, C, alpha=1.0) +``` + +--- + +## FAQ + +### Q: When should I use StaticDispatch vs SizeAutoDispatch? + +Use `StaticDispatch` when the optimal kernel for each shape is **already known and stable**. Use `SizeAutoDispatch` when the best choice depends on runtime factors (micro-architecture quirks, driver versions, etc.) and you want the system to figure it out automatically. + +### Q: What happens if a condition lambda raises an exception? + +The exception propagates uncaught. Keep condition lambdas simple (boolean arithmetic only, no I/O or tensor operations). + +### Q: Can I mix StaticDispatch and SizeAutoDispatch in the same operator? + +Yes. For example, use `StaticDispatch` for well-understood shape classes and `SizeAutoDispatch` for the remainder. Just ensure you return the runner from the appropriate dispatch. + +### Q: Why double-lambda for @libentry kernels? + +`@libentry()` wraps a Triton `JITFunction` in a `LibEntry` object. Calling `entry[grid](args...)` triggers `LibEntry.run()`, which compiles (if needed), launches the kernel, and returns `(kernel_obj, constexprs)`. The inner `lambda:` wraps this into a callable runner without executing the kernel. diff --git a/docs/static_dispatch_cn.md b/docs/static_dispatch_cn.md new file mode 100644 index 00000000..aef4c829 --- /dev/null +++ b/docs/static_dispatch_cn.md @@ -0,0 +1,313 @@ +# StaticDispatch 静态调度表 + +[源码: `src/flag_blas/runtime/dispatch.py`](../src/flag_blas/runtime/dispatch.py) + +`StaticDispatch` 是开发者维护的静态 kernel 调度表,将 shape 条件直接映射到预定的 kernel 工厂函数。**不进行 autotune、不 benchmark、不缓存** —— 条件按顺序求值,首次命中即返回。 + +--- + +## 适用场景 + +当**最优的 shape → kernel 映射已经明确**时(例如通过离线 benchmark、架构分析或领域启发式规则得出),使用 `StaticDispatch`。它提供: + +- **可预测的性能** — 无运行时 autotune 开销 +- **确定性行为** — 相同 shape 总是映射到相同 kernel +- **极简资源占用** — 无缓存文件、无 DB 读写 + +对比 `SizeAutoDispatch`,后者会在运行时对所有候选做 benchmark 并持久化最优选择。 + +--- + +## 核心概念 + +### Condition(条件函数) + +签名:`(m: int, n: int, k: int, aligned: bool, **extra: Any) -> bool` + +判断当前 shape 是否匹配该表条目的 callable。条件**按表中顺序**求值,第一个返回 `True` 的条目胜出。 + +### Factory(工厂函数,双层 Lambda) + +签名:`() -> Callable[[], None]` + +零参数 callable,返回一个 **runner**。runner 本身也是零参数 callable,调用时执行 kernel。 + +**关键要点**:当搭配 `@libentry()` 装饰的 Triton kernel 使用时,factory 必须使用**双层 lambda** 包装: + +```python +# 正确:双层 lambda +lambda: lambda: kernel_fn[grid](arg1, arg2, ...) +# ^^^^ ^^^^^^^^^^^^^^^^^^^^^^^^ +# factory runner(延迟到 runner() 时执行) + +# 错误:单层 lambda —— kernel 会在 factory() 时立即执行 +lambda: kernel_fn[grid](arg1, arg2, ...) +``` + +**原因**:`kernel_fn[grid](args...)` 调用的是 `LibEntry.run()`,它会立即启动 kernel 并返回 `(kernel, constexprs)` 元组,而非可调用的 runner。内层 `lambda:` 将执行延迟到 `runner()` 调用时。 + +--- + +## 架构 + +与 `SizeAutoDispatch` 的多级缓存架构不同,`StaticDispatch` 设计极简: + +``` +lookup_and_build(m, n, k, aligned, **extra) + │ + ├─ 条目 1: condition(m,n,k,aligned)? ─── True → factory() → runner + ├─ 条目 2: condition(m,n,k,aligned)? ─── True → factory() → runner + ├─ ... + └─ 无匹配 → 抛出 ValueError +``` + +- **无过滤逻辑** — 条件内联编码所有匹配规则(无需单独的 `aligned`/`filter` 参数) +- **无缓存** — 每次调用重新求值条件(极其廉价,仅布尔逻辑运算) +- **无匹配即抛错** — 如果没有任何条目匹配,抛出 `ValueError`(最后一条应为兜底条目) + +--- + +## API + +### 构造函数 + +```python +dispatch = StaticDispatch(table) +``` + +| 参数 | 类型 | 说明 | +|------|------|------| +| `table` | `List[Tuple[Condition, Factory]]` | `(condition, factory)` 对的有序列表。**最后一条必须是兜底条目**(condition 始终返回 `True`)。 | + +### lookup_and_build() + +```python +runner = dispatch.lookup_and_build(m, n, k, aligned, **extra) +runner() +``` + +| 参数 | 类型 | 说明 | +|------|------|------| +| `m` | `int` | M 维度 | +| `n` | `int` | N 维度 | +| `k` | `int` | K 维度 | +| `aligned` | `bool` | 输入是否内存对齐 | +| `**extra` | — | 传递给每个 condition 的额外关键字参数 | + +**返回值**:`Callable[[], None]` — 零参数 runner,调用即执行选中的 kernel。 + +**异常**:如果没有任何条目匹配,抛出 `ValueError`(有兜底条目时不应发生)。 + +--- + +## 实战示例:hgemm NN + +来自 [hopper/ops/gemm.py](../src/flag_blas/runtime/backend/_nvidia/hopper/ops/gemm.py) 中 `hgemm` 函数的 NN 分支: + +```python +from flag_blas.runtime.dispatch import StaticDispatch +from triton.tools.tensor_descriptor import TensorDescriptor + +dispatch = StaticDispatch([ + # ── 优先级 1(最高)────────────────────────────────────────── + # Skinny 矩阵 + 对齐 + 大尺寸。 + # 使用 kernel4,搭配 TensorDescriptor 和硬编码最优 config + # (无需 autotune —— 此 config 已证明是该 shape 类型的最优解)。 + ( + lambda m, n, k, aligned, **_kw: + aligned and (m * n > 2048 * 2048) and min(m, n) >= 64 + and ((m >= 16384 and max(n, k) <= 2048) + or (n >= 16384 and max(m, k) <= 2048)), + lambda: lambda: _hgemm_nn_kernel4[( + triton.cdiv(m, 128) * triton.cdiv(n, 256), + )]( + TensorDescriptor( + base=A, shape=[m, k], strides=[lda, 1], + block_shape=[128, 64], + ), + TensorDescriptor( + base=B, shape=[k, n], strides=[ldb, 1], + block_shape=[64, 256], + ), + TensorDescriptor( + base=C, shape=[m, n], strides=[ldc, 1], + block_shape=[128, 256], + ), + alpha, beta, m, n, k, beta_is_zero, + BLOCK_M=128, BLOCK_N=256, BLOCK_K=64, GROUP_M=8, + num_stages=4, num_warps=8, num_ctas=1, + ), + ), + + # ── 优先级 2 ───────────────────────────────────────────────── + # 对齐 + 大尺寸(m×n > 4M,min 维度 ≥ 64)。 + # 使用 kernel3,搭配 TensorDescriptor 和 autotuned configs + # (其 @libtuner 装饰器在运行时选择最优 BLOCK_M/N/K 等参数)。 + ( + lambda m, n, k, aligned, **_kw: + aligned and (m * n > 2048 * 2048) and min(m, n) >= 64, + lambda: lambda: _hgemm_nn_kernel3[grid]( + A, B, C, alpha, beta, m, n, k, + lda, ldb, ldc, beta_is_zero, + ), + ), + + # ── 优先级 3 ───────────────────────────────────────────────── + # 对齐 + 小/中等尺寸(max ≤ 1024)。 + # 使用 kernel2,搭配 block_ptr 和 autotuned configs。 + ( + lambda m, n, k, aligned, **_kw: + aligned and max(m, n) <= 1024, + lambda: lambda: _hgemm_nn_kernel2[grid]( + A, B, C, alpha, beta, m, n, k, + lda, ldb, ldc, beta_is_zero, + ), + ), + + # ── 优先级 4(兜底)────────────────────────────────────────── + # 其余所有情况:未对齐,或中等/大尺寸但未被上述条目覆盖。 + # 使用原始 level3 kernel,基于指针访问。 + ( + lambda **_kw: True, + lambda: lambda: _hgemm_nn_kernel[grid]( + A, B, C, alpha, beta, m, n, k, + lda, ldb, ldc, beta_is_zero, + ), + ), +]) + +runner = dispatch.lookup_and_build(m, n, k, aligned) +runner() +``` + +### 调度逻辑总结 + +| 优先级 | 条件 | Kernel | +|--------|------|--------| +| 1 | 对齐 + 大尺寸 + skinny(单维 ≥ 16384 其他 ≤ 2048) | `kernel4` — TensorDescriptor,硬编码 config | +| 2 | 对齐 + 大尺寸(m×n > 2048²,min ≥ 64) | `kernel3` — TensorDescriptor,autotuned | +| 3 | 对齐 + 小尺寸(max ≤ 1024) | `kernel2` — block_ptr,autotuned | +| 4 | 其余所有 | `kernel` — 原始指针实现,autotuned | + +### Kernel 变体特征 + +| Kernel | 来源 | 数据访问 | Config 策略 | 最佳场景 | +|--------|------|---------|-------------|----------| +| `_hgemm_nn_kernel4` | Hopper(本地) | `TensorDescriptor.load/store` + `int32` 偏移 | 硬编码 `(128,256,64,8,4,8,1)` | Skinny 大矩阵 | +| `_hgemm_nn_kernel3` | Hopper(本地) | `TensorDescriptor.load/store` | `@libtuner("hgemm_nn2")` | 常规大矩阵 | +| `_hgemm_nn_kernel2` | Hopper(本地) | `tl.make_block_ptr` + `tl.advance` | `@libtuner("hgemm_nn")` | 小/中等对齐 | +| `_hgemm_nn_kernel` | Level3(导入) | 原始指针 + `offs_{am,bn,k}` + mask 逻辑 | `@libtuner("hgemm_nn")` | 其余情况 | + +--- + +## 设计原则 + +1. **无 autotune**:表由人工维护;无运行时 benchmark。 +2. **顺序匹配**:条件自上而下求值;首次 `True` 即胜出。 +3. **必须兜底**:最后一条必须匹配所有 shape(防止 `ValueError`)。 +4. **条件互斥**:条目间不应重叠以保证行为可预测。高优先级条目应有更具体的条件。 +5. **双层 lambda 工厂**:搭配 `@libentry()` 装饰的 Triton kernel 时必须使用。内层 `lambda:` 将 `LibEntry.run()` 延迟到 `runner()` 时。 + +--- + +## 与 SizeAutoDispatch 对比 + +| | `SizeAutoDispatch` | `StaticDispatch` | +|---|---|---| +| **选择方式** | Autotune(benchmark 全部 → 选最快) | 静态条件(首次命中即返回) | +| **缓存** | 内存 + SQLite DB | 无 | +| **持久化** | 跨进程,SQLite | 无 | +| **首次调用** | benchmark 开销(秒级) | 瞬时(仅条件求值) | +| **后续调用** | ~缓存查找 | 瞬时(同首次调用) | +| **表构建** | 每个变体 `add()` | 构造函数一次性传入列表 | +| **失败处理** | 降级到第一个候选 | `ValueError`(兜底条目规避) | +| **最佳用途** | 最优映射未知 | 最优映射已知 | + +--- + +## 构建 StaticDispatch 表 + +### 步骤 + +1. **列出 kernel 变体**。列出所有能处理该运算的 kernel,注明各自优势。 + +2. **编写条件函数**。为每个变体定义 lambda,在其表现最优的 shape 上返回 `True`。 + +3. **按优先级排序**。最具体的条件(最窄范围)放在最前面,最宽泛的放在最后。 + +4. **添加兜底条目**。最后一条必须匹配所有情况(`lambda **_kw: True`)。 + +5. **测试边界**。确保处于条件边界的 shape 路由到预期的 kernel。开发时可用日志或调试打印来验证。 + +### 常见模式 + +**按维度分优先级**: +```python +[ + (lambda m, *_kw, **__: m > 8192, factory_a), # 超大 + (lambda m, *_kw, **__: m > 1024, factory_b), # 大 + (lambda m, *_kw, **__: m > 256, factory_c), # 中等 + (lambda **_kw: True, factory_d), # 小 +] +``` + +**按对齐分优先级**: +```python +[ + (lambda aligned, **_kw: aligned and is_large(**kw), aligned_large_factory), + (lambda aligned, **_kw: aligned, aligned_small_factory), + (lambda **_kw: True, unaligned_factory), +] +``` + +**维度 + 对齐组合**(如 hgemm_nn): +```python +[ + (lambda aligned, m, n, k, **_kw: + aligned and meets_criteria_A(m, n, k), + factory_a), + (lambda aligned, m, n, k, **_kw: + aligned and meets_criteria_B(m, n, k), + factory_b), + (lambda **_kw: True, + fallback_factory), +] +``` + +--- + +## 辅助类:KernelRunner + +```python +class KernelRunner: + def __init__(self, kernel: Callable, *args, **kwargs): ... + def __call__(self): + return self._kernel(*self._args, **self._kwargs) +``` + +将 kernel 函数与其参数绑定的简单可调用包装器。适用于不需要 autotune、只需固定实现的场景: + +```python +runner = KernelRunner(my_kernel_fn, A, B, C, alpha=1.0) +runner() # 等价于 my_kernel_fn(A, B, C, alpha=1.0) +``` + +--- + +## 常见问题 + +### Q: 什么时候应该用 StaticDispatch 而不是 SizeAutoDispatch? + +当每个 shape 对应的最优 kernel **已知且稳定**时用 `StaticDispatch`。当最优选择依赖运行时因素(微架构特性、驱动版本等)且你希望系统自动发现时,用 `SizeAutoDispatch`。 + +### Q: 如果条件 lambda 抛出异常会怎样? + +异常会向上传播不被捕获。保持条件 lambda 尽量简单(仅布尔算术,无 IO 或 tensor 操作)。 + +### Q: 能在同一个算子里混用 StaticDispatch 和 SizeAutoDispatch 吗? + +可以。例如,用 `StaticDispatch` 处理已知最优的 shape 类型,用 `SizeAutoDispatch` 处理其余情况。只需确保从正确的 dispatch 返回 runner 即可。 + +### Q: 为什么 @libentry kernel 需要双层 lambda? + +`@libentry()` 将 Triton `JITFunction` 包装为 `LibEntry` 对象。调用 `entry[grid](args...)` 会触发 `LibEntry.run()`,该方法编译(如果需要)、启动 kernel 并返回 `(kernel_obj, constexprs)`。内层 `lambda:` 将这个过程包装为可调用的 runner 而不立即执行 kernel。 From 73e1e77b1c7bd12b7104e54feed9e3b43e3594d3 Mon Sep 17 00:00:00 2001 From: bin913 <842884726@qq.com> Date: Mon, 22 Jun 2026 16:05:03 +0800 Subject: [PATCH 5/6] optimize static dispatch --- .../backend/_nvidia/hopper/ops/gemm.py | 138 ++++++++++++------ src/flag_blas/runtime/dispatch.py | 4 + 2 files changed, 101 insertions(+), 41 deletions(-) diff --git a/src/flag_blas/runtime/backend/_nvidia/hopper/ops/gemm.py b/src/flag_blas/runtime/backend/_nvidia/hopper/ops/gemm.py index 3033e291..8230d7b0 100644 --- a/src/flag_blas/runtime/backend/_nvidia/hopper/ops/gemm.py +++ b/src/flag_blas/runtime/backend/_nvidia/hopper/ops/gemm.py @@ -1800,6 +1800,96 @@ def _hgemm_tt_kernel3( desc_c.store([offs_m, offs_n], result) +# --------------------------------------------------------------------------- +# Module-level condition predicates for hgemm StaticDispatch +# --------------------------------------------------------------------------- +def _hgemm_nn_is_skinny_aligned_large(m, n, k, aligned, **_kw): + return ( + aligned + and (m * n > 2048 * 2048) + and min(m, n) >= 64 + and ( + (m >= 16384 and max(n, k) <= 2048) + or (n >= 16384 and max(m, k) <= 2048) + ) + ) + + +def _hgemm_nn_is_aligned_large(m, n, k, aligned, **_kw): + return aligned and (m * n > 2048 * 2048) and min(m, n) >= 64 + + +def _hgemm_nn_is_aligned_small(m, n, k, aligned, **_kw): + return aligned and max(m, n) <= 1024 + + +def _hgemm_nn_is_default(**_kw): + return True + + +# --------------------------------------------------------------------------- +# Module-level factory functions for hgemm StaticDispatch +# --------------------------------------------------------------------------- +def _hgemm_nn_build_kernel4( + A, B, C, m, n, k, lda, ldb, ldc, alpha, beta, beta_is_zero, +): + return lambda: _hgemm_nn_kernel4[( + triton.cdiv(m, 128) * triton.cdiv(n, 256), + )]( + TensorDescriptor(base=A, shape=[m, k], strides=[lda, 1], block_shape=[128, 64]), + TensorDescriptor(base=B, shape=[k, n], strides=[ldb, 1], block_shape=[64, 256]), + TensorDescriptor(base=C, shape=[m, n], strides=[ldc, 1], block_shape=[128, 256]), + alpha, beta, m, n, k, beta_is_zero, + BLOCK_M=128, BLOCK_N=256, BLOCK_K=64, GROUP_M=8, + num_stages=4, num_warps=8, num_ctas=1, + ) + + +def _hgemm_nn_build_kernel3( + A, B, C, m, n, k, lda, ldb, ldc, alpha, beta, beta_is_zero, +): + grid = lambda meta: ( + triton.cdiv(m, meta["BLOCK_M"]) * triton.cdiv(n, meta["BLOCK_N"]), + ) + return lambda: _hgemm_nn_kernel3[grid]( + A, B, C, alpha, beta, m, n, k, lda, ldb, ldc, beta_is_zero, + ) + + +def _hgemm_nn_build_kernel2( + A, B, C, m, n, k, lda, ldb, ldc, alpha, beta, beta_is_zero, +): + grid = lambda meta: ( + triton.cdiv(m, meta["BLOCK_M"]) * triton.cdiv(n, meta["BLOCK_N"]), + ) + return lambda: _hgemm_nn_kernel2[grid]( + A, B, C, alpha, beta, m, n, k, lda, ldb, ldc, beta_is_zero, + ) + + +def _hgemm_nn_build_kernel( + A, B, C, m, n, k, lda, ldb, ldc, alpha, beta, beta_is_zero, +): + grid = lambda meta: ( + triton.cdiv(m, meta["BLOCK_M"]) * triton.cdiv(n, meta["BLOCK_N"]), + ) + return lambda: _hgemm_nn_kernel[grid]( + A, B, C, alpha, beta, m, n, k, lda, ldb, ldc, beta_is_zero, + ) + + +_HGEMM_NN_DISPATCH = StaticDispatch([ + # skinny + aligned large → kernel4 (TensorDescriptor, hardcoded config) + (_hgemm_nn_is_skinny_aligned_large, _hgemm_nn_build_kernel4), + # aligned large → kernel3 (TensorDescriptor, autotuned config) + (_hgemm_nn_is_aligned_large, _hgemm_nn_build_kernel3), + # aligned small → kernel2 (block_ptr) + (_hgemm_nn_is_aligned_small, _hgemm_nn_build_kernel2), + # default → kernel (original) + (_hgemm_nn_is_default, _hgemm_nn_build_kernel), +]) + + def hgemm( transa: int, transb: int, @@ -1861,48 +1951,14 @@ def hgemm( aligned = _is_gemm_aligned(A, lda, B, ldb, C, ldc) with torch_device_fn.device(A.device): if transa == CUBLAS_OP_N and transb == CUBLAS_OP_N: - dispatch = StaticDispatch([ - # skinny + aligned large → kernel4 (TensorDescriptor, hardcoded config) - ( - lambda m, n, k, aligned, **_kw: - aligned and (m * n > 2048 * 2048) and min(m, n) >= 64 - and ((m >= 16384 and max(n, k) <= 2048) or (n >= 16384 and max(m, k) <= 2048)), - lambda: lambda: _hgemm_nn_kernel4[( - triton.cdiv(m, 128) * triton.cdiv(n, 256), - )]( - TensorDescriptor(base=A, shape=[m, k], strides=[lda, 1], block_shape=[128, 64]), - TensorDescriptor(base=B, shape=[k, n], strides=[ldb, 1], block_shape=[64, 256]), - TensorDescriptor(base=C, shape=[m, n], strides=[ldc, 1], block_shape=[128, 256]), - alpha, beta, m, n, k, beta_is_zero, - BLOCK_M=128, BLOCK_N=256, BLOCK_K=64, GROUP_M=8, - num_stages=4, num_warps=8, num_ctas=1, - ), - ), - # aligned large → kernel3 (TensorDescriptor, autotuned config) - ( - lambda m, n, k, aligned, **_kw: - aligned and (m * n > 2048 * 2048) and min(m, n) >= 64, - lambda: lambda: _hgemm_nn_kernel3[grid]( - A, B, C, alpha, beta, m, n, k, lda, ldb, ldc, beta_is_zero, - ), - ), - # aligned small → kernel2 (block_ptr) - ( - lambda m, n, k, aligned, **_kw: - aligned and max(m, n) <= 1024, - lambda: lambda: _hgemm_nn_kernel2[grid]( - A, B, C, alpha, beta, m, n, k, lda, ldb, ldc, beta_is_zero, - ), + runner = _HGEMM_NN_DISPATCH.lookup_and_build( + m, n, k, aligned, + context=dict( + A=A, B=B, C=C, m=m, n=n, k=k, + lda=lda, ldb=ldb, ldc=ldc, + alpha=alpha, beta=beta, beta_is_zero=beta_is_zero, ), - # default → kernel (original) - ( - lambda **_kw: True, - lambda: lambda: _hgemm_nn_kernel[grid]( - A, B, C, alpha, beta, m, n, k, lda, ldb, ldc, beta_is_zero, - ), - ), - ]) - runner = dispatch.lookup_and_build(m, n, k, aligned) + ) runner() elif transa == CUBLAS_OP_T and transb == CUBLAS_OP_N: is_skinny = (m >= 16384 and max(n, k) <= 2048) or ( diff --git a/src/flag_blas/runtime/dispatch.py b/src/flag_blas/runtime/dispatch.py index 1a1991f0..2d57a95e 100644 --- a/src/flag_blas/runtime/dispatch.py +++ b/src/flag_blas/runtime/dispatch.py @@ -367,10 +367,14 @@ def lookup_and_build( n: int, k: int, aligned: bool, + *, + context: Optional[dict] = None, **extra, ) -> Callable[[], None]: for condition, factory in self._table: if condition(m=m, n=n, k=k, aligned=aligned, **extra): + if context is not None: + return factory(**context) return factory() raise ValueError( f"StaticDispatch: no matching entry for " From a139bca37e6ce502e21d68850f97fc18cc6cf63f Mon Sep 17 00:00:00 2001 From: bin913 <842884726@qq.com> Date: Tue, 23 Jun 2026 11:51:23 +0800 Subject: [PATCH 6/6] fix docs for static_dispatch --- docs/static_dispatch.md | 298 +++++++++++++++++++----------- docs/static_dispatch_cn.md | 298 +++++++++++++++++++----------- src/flag_blas/runtime/dispatch.py | 47 ++++- 3 files changed, 420 insertions(+), 223 deletions(-) diff --git a/docs/static_dispatch.md b/docs/static_dispatch.md index 7c1e10bd..31d55450 100644 --- a/docs/static_dispatch.md +++ b/docs/static_dispatch.md @@ -4,6 +4,8 @@ `StaticDispatch` is a developer-maintained static dispatch table that maps shape conditions directly to pre-determined kernel factories. **No autotune, no benchmarking, no caching** — conditions are evaluated in order, and the first match wins immediately. +The dispatch table is designed to be created once at module level and reused across calls. Per-call varying data (tensors, scalars) is passed through a `context` dict, so factories don't need to capture variables via closures. + --- ## When to Use @@ -26,25 +28,37 @@ Signature: `(m: int, n: int, k: int, aligned: bool, **extra: Any) -> bool` A callable that determines whether the current shape matches this table entry. Conditions are evaluated **in table order**; the first that returns `True` wins. -### Factory (Double-Lambda) +### Factory -Signature: `() -> Callable[[], None]` +Signature: `(context_key1=..., context_key2=..., ...) -> Callable[[], None]` -A zero-arg callable that returns a **runner**. The runner is itself a zero-arg callable that executes the kernel. +A named function (recommended over lambdas) that accepts per-call varying arguments via keyword arguments and returns a **runner** — a zero-arg callable that executes the kernel. The arguments are passed from the `context` dict in `lookup_and_build()`. -**Critical**: When used with `@libentry()`-decorated Triton kernels, the factory must use a **double-lambda** wrapper: +When the dispatch table is created at module level, factories are pure references to named functions — no closures, no lambdas rebuilt on every call. The per-call data flows in through `context`. -```python -# Correct: double lambda -lambda: lambda: kernel_fn[grid](arg1, arg2, ...) -# ^^^^ ^^^^^^^^^^^^^^^^^^^^^^^^ -# factory runner (deferred to runner()) +**Critical**: When used with `@libentry()`-decorated Triton kernels, the factory must return the kernel call wrapped in a **lambda**: -# Wrong: single lambda — kernel executes immediately in factory() -lambda: kernel_fn[grid](arg1, arg2, ...) +```python +# Module level — defined once: +def build_my_kernel(A, B, C, alpha, beta, m, n, k, lda, ldb, ldc): + grid = lambda meta: ( + triton.cdiv(m, meta["BLOCK_M"]) * triton.cdiv(n, meta["BLOCK_N"]), + ) + return lambda: _my_kernel[grid]( + A, B, C, alpha, beta, m, n, k, lda, ldb, ldc, + ) + +# Per-call — context carries the varying data: +runner = dispatch.lookup_and_build( + m, n, k, aligned, + context=dict(A=A, B=B, C=C, m=m, n=n, k=k, + lda=lda, ldb=ldb, ldc=ldc, + alpha=alpha, beta=beta), +) +runner() ``` -**Why**: `kernel_fn[grid](args...)` calls `LibEntry.run()`, which launches the kernel immediately and returns a `(kernel, constexprs)` tuple — not a callable runner. The extra `lambda:` defers execution to `runner()` time. +**Why lambda**: `kernel_fn[grid](args...)` calls `LibEntry.run()`, which launches the kernel immediately and returns a `(kernel, constexprs)` tuple — not a callable runner. The `lambda:` defers execution to `runner()` time. --- @@ -53,16 +67,17 @@ lambda: kernel_fn[grid](arg1, arg2, ...) Unlike `SizeAutoDispatch` with its multi-tier cache, `StaticDispatch` has a minimal design: ``` -lookup_and_build(m, n, k, aligned, **extra) +lookup_and_build(m, n, k, aligned, *, context, **extra) │ - ├─ Entry 1: condition(m,n,k,aligned)? ─── True → factory() → runner - ├─ Entry 2: condition(m,n,k,aligned)? ─── True → factory() → runner + ├─ Entry 1: condition(m,n,k,aligned)? ─── True → factory(**context) → runner + ├─ Entry 2: condition(m,n,k,aligned)? ─── True → factory(**context) → runner ├─ ... └─ No match → raise ValueError ``` - **No filtering logic** — conditions encode all matching criteria inline (no separate `aligned`/`filter` params) - **No cache** — every call re-evaluates conditions (cheap, just boolean logic) +- **`context` dict** — passes per-call varying data (tensors, scalars) to factories, so the dispatch table itself can live at module level - **Throw on miss** — if no entry matches, raises `ValueError` (the last entry should always be a catch-all) --- @@ -82,7 +97,7 @@ dispatch = StaticDispatch(table) ### lookup_and_build() ```python -runner = dispatch.lookup_and_build(m, n, k, aligned, **extra) +runner = dispatch.lookup_and_build(m, n, k, aligned, *, context=None, **extra) runner() ``` @@ -92,6 +107,7 @@ runner() | `n` | `int` | N dimension | | `k` | `int` | K dimension | | `aligned` | `bool` | Whether inputs are memory-aligned | +| `context` | `dict` or `None` | Per-call varying data (tensors, scalars, etc.) passed as keyword arguments to the matched factory. When `None`, factories are called with no arguments. | | `**extra` | — | Additional keyword arguments passed to each condition | **Returns**: `Callable[[], None]` — zero-arg runner; calling it executes the selected kernel. @@ -102,82 +118,119 @@ runner() ## Real-World Example: hgemm NN -This example comes from [hopper/ops/gemm.py](../src/flag_blas/runtime/backend/_nvidia/hopper/ops/gemm.py), the `hgemm` function's NN branch: +This example comes from [hopper/ops/gemm.py](../src/flag_blas/runtime/backend/_nvidia/hopper/ops/gemm.py), the `hgemm` function's NN branch. + +### Module level — defined once ```python from flag_blas.runtime.dispatch import StaticDispatch from triton.tools.tensor_descriptor import TensorDescriptor -dispatch = StaticDispatch([ - # ── Priority 1 (highest) ───────────────────────────────────── - # Skinny matrix with aligned large dimensions. - # Uses kernel4 with TensorDescriptor and hardcoded optimal config - # (no autotune needed — config proven optimal for this shape class). - ( - lambda m, n, k, aligned, **_kw: - aligned and (m * n > 2048 * 2048) and min(m, n) >= 64 - and ((m >= 16384 and max(n, k) <= 2048) - or (n >= 16384 and max(m, k) <= 2048)), - lambda: lambda: _hgemm_nn_kernel4[( - triton.cdiv(m, 128) * triton.cdiv(n, 256), - )]( - TensorDescriptor( - base=A, shape=[m, k], strides=[lda, 1], - block_shape=[128, 64], - ), - TensorDescriptor( - base=B, shape=[k, n], strides=[ldb, 1], - block_shape=[64, 256], - ), - TensorDescriptor( - base=C, shape=[m, n], strides=[ldc, 1], - block_shape=[128, 256], - ), - alpha, beta, m, n, k, beta_is_zero, - BLOCK_M=128, BLOCK_N=256, BLOCK_K=64, GROUP_M=8, - num_stages=4, num_warps=8, num_ctas=1, - ), - ), - - # ── Priority 2 ─────────────────────────────────────────────── - # Aligned + large dimensions (m×n > 4M, min dim ≥ 64). - # Uses kernel3 with TensorDescriptor and autotuned configs - # (its @libtuner picks the best BLOCK_M/N/K etc. at runtime). - ( - lambda m, n, k, aligned, **_kw: - aligned and (m * n > 2048 * 2048) and min(m, n) >= 64, - lambda: lambda: _hgemm_nn_kernel3[grid]( - A, B, C, alpha, beta, m, n, k, - lda, ldb, ldc, beta_is_zero, +# ── Condition predicates (named functions, not lambdas) ────────── + +def _hgemm_nn_is_skinny_aligned_large(m, n, k, aligned, **_kw): + return ( + aligned + and (m * n > 2048 * 2048) + and min(m, n) >= 64 + and ( + (m >= 16384 and max(n, k) <= 2048) + or (n >= 16384 and max(m, k) <= 2048) + ) + ) + +def _hgemm_nn_is_aligned_large(m, n, k, aligned, **_kw): + return aligned and (m * n > 2048 * 2048) and min(m, n) >= 64 + +def _hgemm_nn_is_aligned_small(m, n, k, aligned, **_kw): + return aligned and max(m, n) <= 1024 + +def _hgemm_nn_is_default(**_kw): + return True + +# ── Factory functions (accept context dict keys as kwargs) ─────── + +def _hgemm_nn_build_kernel4(A, B, C, m, n, k, lda, ldb, ldc, + alpha, beta, beta_is_zero): + return lambda: _hgemm_nn_kernel4[( + triton.cdiv(m, 128) * triton.cdiv(n, 256), + )]( + TensorDescriptor( + base=A, shape=[m, k], strides=[lda, 1], + block_shape=[128, 64], ), - ), - - # ── Priority 3 ─────────────────────────────────────────────── - # Aligned + small/moderate dimensions (max ≤ 1024). - # Uses kernel2 with block_ptr and autotuned configs. - ( - lambda m, n, k, aligned, **_kw: - aligned and max(m, n) <= 1024, - lambda: lambda: _hgemm_nn_kernel2[grid]( - A, B, C, alpha, beta, m, n, k, - lda, ldb, ldc, beta_is_zero, + TensorDescriptor( + base=B, shape=[k, n], strides=[ldb, 1], + block_shape=[64, 256], ), - ), - - # ── Priority 4 (catch-all) ─────────────────────────────────── - # Everything else: unaligned or moderate/large but not covered above. - # Uses the original level3 kernel with pointer-based access. - ( - lambda **_kw: True, - lambda: lambda: _hgemm_nn_kernel[grid]( - A, B, C, alpha, beta, m, n, k, - lda, ldb, ldc, beta_is_zero, + TensorDescriptor( + base=C, shape=[m, n], strides=[ldc, 1], + block_shape=[128, 256], ), - ), + alpha, beta, m, n, k, beta_is_zero, + BLOCK_M=128, BLOCK_N=256, BLOCK_K=64, GROUP_M=8, + num_stages=4, num_warps=8, num_ctas=1, + ) + +def _hgemm_nn_build_kernel3(A, B, C, m, n, k, lda, ldb, ldc, + alpha, beta, beta_is_zero): + grid = lambda meta: ( + triton.cdiv(m, meta["BLOCK_M"]) * triton.cdiv(n, meta["BLOCK_N"]), + ) + return lambda: _hgemm_nn_kernel3[grid]( + A, B, C, alpha, beta, m, n, k, lda, ldb, ldc, beta_is_zero, + ) + +def _hgemm_nn_build_kernel2(A, B, C, m, n, k, lda, ldb, ldc, + alpha, beta, beta_is_zero): + grid = lambda meta: ( + triton.cdiv(m, meta["BLOCK_M"]) * triton.cdiv(n, meta["BLOCK_N"]), + ) + return lambda: _hgemm_nn_kernel2[grid]( + A, B, C, alpha, beta, m, n, k, lda, ldb, ldc, beta_is_zero, + ) + +def _hgemm_nn_build_kernel(A, B, C, m, n, k, lda, ldb, ldc, + alpha, beta, beta_is_zero): + grid = lambda meta: ( + triton.cdiv(m, meta["BLOCK_M"]) * triton.cdiv(n, meta["BLOCK_N"]), + ) + return lambda: _hgemm_nn_kernel[grid]( + A, B, C, alpha, beta, m, n, k, lda, ldb, ldc, beta_is_zero, + ) + +_HGEMM_NN_DISPATCH = StaticDispatch([ + # skinny + aligned large → kernel4 (TensorDescriptor, hardcoded config) + (_hgemm_nn_is_skinny_aligned_large, _hgemm_nn_build_kernel4), + # aligned large → kernel3 (TensorDescriptor, autotuned config) + (_hgemm_nn_is_aligned_large, _hgemm_nn_build_kernel3), + # aligned small → kernel2 (block_ptr) + (_hgemm_nn_is_aligned_small, _hgemm_nn_build_kernel2), + # default → kernel (original) + (_hgemm_nn_is_default, _hgemm_nn_build_kernel), ]) +``` -runner = dispatch.lookup_and_build(m, n, k, aligned) -runner() +### Per-call — inside `hgemm()` + +```python +def hgemm(transa, transb, m, n, k, alpha, A, lda, B, ldb, beta, C, ldc): + # ... validation ... + beta_is_zero = beta == 0.0 + aligned = _is_gemm_aligned(A, lda, B, ldb, C, ldc) + + with torch_device_fn.device(A.device): + if transa == CUBLAS_OP_N and transb == CUBLAS_OP_N: + runner = _HGEMM_NN_DISPATCH.lookup_and_build( + m, n, k, aligned, + context=dict( + A=A, B=B, C=C, m=m, n=n, k=k, + lda=lda, ldb=ldb, ldc=ldc, + alpha=alpha, beta=beta, beta_is_zero=beta_is_zero, + ), + ) + runner() + # ... other transa/transb branches ... ``` ### Dispatch Logic Summary @@ -206,7 +259,7 @@ runner() 2. **Ordered matching**: Conditions evaluated top-to-bottom; first `True` wins. 3. **Catch-all required**: The last entry must match any shape (prevent `ValueError`). 4. **Mutually exclusive conditions**: Entries should not overlap to make behavior predictable. Higher-priority entries should have more specific conditions. -5. **Double-lambda factories**: Required when using `@libentry()`-decorated Triton kernels. The inner `lambda:` defers `LibEntry.run()` to `runner()` time. +5. **Named functions, not lambdas in the table**: Conditions and factories should be module-level named functions referenced by name in the `StaticDispatch` table. Per-call varying data (tensors, scalars) is passed via the `context` dict to `lookup_and_build()`. This avoids recreating closures on every call. --- @@ -241,37 +294,68 @@ runner() ### Common Patterns -**Priority by dimension**: +**Priority by dimension** (module-level named functions): ```python -[ - (lambda m, *_kw, **__: m > 8192, factory_a), # very large - (lambda m, *_kw, **__: m > 1024, factory_b), # large - (lambda m, *_kw, **__: m > 256, factory_c), # medium - (lambda **_kw: True, factory_d), # small -] +def is_very_large(m, **_kw): + return m > 8192 + +def is_large(m, **_kw): + return m > 1024 + +def is_medium(m, **_kw): + return m > 256 + +def is_default(**_kw): + return True + +_DISPATCH = StaticDispatch([ + (is_very_large, build_kernel_a), + (is_large, build_kernel_b), + (is_medium, build_kernel_c), + (is_default, build_kernel_d), +]) ``` **Priority by alignment**: ```python -[ - (lambda aligned, **_kw: aligned and is_large(**kw), aligned_large_factory), - (lambda aligned, **_kw: aligned, aligned_small_factory), - (lambda **_kw: True, unaligned_factory), -] +def is_aligned_large(aligned, m, n, k, **_kw): + return aligned and (m * n > 2048 * 2048) + +def is_aligned_only(aligned, **_kw): + return aligned + +def is_default(**_kw): + return True + +_DISPATCH = StaticDispatch([ + (is_aligned_large, build_aligned_large), + (is_aligned_only, build_aligned), + (is_default, build_fallback), +]) ``` **Combined dimensions + alignment** (as in hgemm_nn): ```python -[ - (lambda aligned, m, n, k, **_kw: - aligned and meets_criteria_A(m, n, k), - factory_a), - (lambda aligned, m, n, k, **_kw: - aligned and meets_criteria_B(m, n, k), - factory_b), - (lambda **_kw: True, - fallback_factory), -] +def is_skinny_aligned_large(m, n, k, aligned, **_kw): + return (aligned and (m * n > 2048 * 2048) and min(m, n) >= 64 + and ((m >= 16384 and max(n, k) <= 2048) + or (n >= 16384 and max(m, k) <= 2048))) + +def is_aligned_large(m, n, k, aligned, **_kw): + return aligned and (m * n > 2048 * 2048) and min(m, n) >= 64 + +def is_aligned_small(m, n, k, aligned, **_kw): + return aligned and max(m, n) <= 1024 + +def is_default(**_kw): + return True + +_DISPATCH = StaticDispatch([ + (is_skinny_aligned_large, build_kernel4), + (is_aligned_large, build_kernel3), + (is_aligned_small, build_kernel2), + (is_default, build_kernel), +]) ``` --- @@ -308,6 +392,6 @@ The exception propagates uncaught. Keep condition lambdas simple (boolean arithm Yes. For example, use `StaticDispatch` for well-understood shape classes and `SizeAutoDispatch` for the remainder. Just ensure you return the runner from the appropriate dispatch. -### Q: Why double-lambda for @libentry kernels? +### Q: Why return a lambda from the factory? -`@libentry()` wraps a Triton `JITFunction` in a `LibEntry` object. Calling `entry[grid](args...)` triggers `LibEntry.run()`, which compiles (if needed), launches the kernel, and returns `(kernel_obj, constexprs)`. The inner `lambda:` wraps this into a callable runner without executing the kernel. +`@libentry()` wraps a Triton `JITFunction` in a `LibEntry` object. Calling `entry[grid](args...)` triggers `LibEntry.run()`, which compiles (if needed), launches the kernel, and returns `(kernel_obj, constexprs)`. The returned `lambda:` wraps this into a callable runner without executing the kernel immediately. diff --git a/docs/static_dispatch_cn.md b/docs/static_dispatch_cn.md index aef4c829..cfccfeab 100644 --- a/docs/static_dispatch_cn.md +++ b/docs/static_dispatch_cn.md @@ -4,6 +4,8 @@ `StaticDispatch` 是开发者维护的静态 kernel 调度表,将 shape 条件直接映射到预定的 kernel 工厂函数。**不进行 autotune、不 benchmark、不缓存** —— 条件按顺序求值,首次命中即返回。 +调度表设计为模块级别创建一次,跨调用复用。每次调用的变化数据(tensor、标量等)通过 `context` 字典传入 `lookup_and_build()`,工厂函数无需通过闭包捕获变量。 + --- ## 适用场景 @@ -26,25 +28,37 @@ 判断当前 shape 是否匹配该表条目的 callable。条件**按表中顺序**求值,第一个返回 `True` 的条目胜出。 -### Factory(工厂函数,双层 Lambda) +### Factory(工厂函数) -签名:`() -> Callable[[], None]` +签名:`(context_key1=..., context_key2=..., ...) -> Callable[[], None]` -零参数 callable,返回一个 **runner**。runner 本身也是零参数 callable,调用时执行 kernel。 +一个命名函数(推荐使用命名函数而非 lambda),通过关键字参数接收每次调用的变化数据,返回一个 **runner** —— 执行 kernel 的零参数可调用对象。参数通过 `lookup_and_build()` 中的 `context` 字典传入。 -**关键要点**:当搭配 `@libentry()` 装饰的 Triton kernel 使用时,factory 必须使用**双层 lambda** 包装: +当调度表在模块级别创建时,工厂函数只是对命名函数的纯引用 —— 没有闭包,没有每次调用重新构建的 lambda。每次调用的数据通过 `context` 流入。 -```python -# 正确:双层 lambda -lambda: lambda: kernel_fn[grid](arg1, arg2, ...) -# ^^^^ ^^^^^^^^^^^^^^^^^^^^^^^^ -# factory runner(延迟到 runner() 时执行) +**关键要点**:当搭配 `@libentry()` 装饰的 Triton kernel 使用时,工厂必须返回一个被 **lambda** 包装的 kernel 调用: -# 错误:单层 lambda —— kernel 会在 factory() 时立即执行 -lambda: kernel_fn[grid](arg1, arg2, ...) +```python +# 模块级别 —— 定义一次: +def build_my_kernel(A, B, C, alpha, beta, m, n, k, lda, ldb, ldc): + grid = lambda meta: ( + triton.cdiv(m, meta["BLOCK_M"]) * triton.cdiv(n, meta["BLOCK_N"]), + ) + return lambda: _my_kernel[grid]( + A, B, C, alpha, beta, m, n, k, lda, ldb, ldc, + ) + +# 每次调用 —— context 携带变化数据: +runner = dispatch.lookup_and_build( + m, n, k, aligned, + context=dict(A=A, B=B, C=C, m=m, n=n, k=k, + lda=lda, ldb=ldb, ldc=ldc, + alpha=alpha, beta=beta), +) +runner() ``` -**原因**:`kernel_fn[grid](args...)` 调用的是 `LibEntry.run()`,它会立即启动 kernel 并返回 `(kernel, constexprs)` 元组,而非可调用的 runner。内层 `lambda:` 将执行延迟到 `runner()` 调用时。 +**原因**:`kernel_fn[grid](args...)` 调用的是 `LibEntry.run()`,它会立即启动 kernel 并返回 `(kernel_obj, constexprs)` 元组,而非可调用的 runner。`lambda:` 将执行延迟到 `runner()` 调用时。 --- @@ -53,16 +67,17 @@ lambda: kernel_fn[grid](arg1, arg2, ...) 与 `SizeAutoDispatch` 的多级缓存架构不同,`StaticDispatch` 设计极简: ``` -lookup_and_build(m, n, k, aligned, **extra) +lookup_and_build(m, n, k, aligned, *, context, **extra) │ - ├─ 条目 1: condition(m,n,k,aligned)? ─── True → factory() → runner - ├─ 条目 2: condition(m,n,k,aligned)? ─── True → factory() → runner + ├─ 条目 1: condition(m,n,k,aligned)? ─── True → factory(**context) → runner + ├─ 条目 2: condition(m,n,k,aligned)? ─── True → factory(**context) → runner ├─ ... └─ 无匹配 → 抛出 ValueError ``` - **无过滤逻辑** — 条件内联编码所有匹配规则(无需单独的 `aligned`/`filter` 参数) - **无缓存** — 每次调用重新求值条件(极其廉价,仅布尔逻辑运算) +- **`context` 字典** — 将每次调用的变化数据(tensor、标量等)传入工厂函数,使调度表本身可以驻留在模块级别 - **无匹配即抛错** — 如果没有任何条目匹配,抛出 `ValueError`(最后一条应为兜底条目) --- @@ -82,7 +97,7 @@ dispatch = StaticDispatch(table) ### lookup_and_build() ```python -runner = dispatch.lookup_and_build(m, n, k, aligned, **extra) +runner = dispatch.lookup_and_build(m, n, k, aligned, *, context=None, **extra) runner() ``` @@ -92,6 +107,7 @@ runner() | `n` | `int` | N 维度 | | `k` | `int` | K 维度 | | `aligned` | `bool` | 输入是否内存对齐 | +| `context` | `dict` 或 `None` | 每次调用的变化数据(tensor、标量等),以关键字参数形式传入匹配的工厂函数。为 `None` 时工厂函数无参调用。 | | `**extra` | — | 传递给每个 condition 的额外关键字参数 | **返回值**:`Callable[[], None]` — 零参数 runner,调用即执行选中的 kernel。 @@ -102,82 +118,119 @@ runner() ## 实战示例:hgemm NN -来自 [hopper/ops/gemm.py](../src/flag_blas/runtime/backend/_nvidia/hopper/ops/gemm.py) 中 `hgemm` 函数的 NN 分支: +来自 [hopper/ops/gemm.py](../src/flag_blas/runtime/backend/_nvidia/hopper/ops/gemm.py) 中 `hgemm` 函数的 NN 分支。 + +### 模块级别 —— 定义一次 ```python from flag_blas.runtime.dispatch import StaticDispatch from triton.tools.tensor_descriptor import TensorDescriptor -dispatch = StaticDispatch([ - # ── 优先级 1(最高)────────────────────────────────────────── - # Skinny 矩阵 + 对齐 + 大尺寸。 - # 使用 kernel4,搭配 TensorDescriptor 和硬编码最优 config - # (无需 autotune —— 此 config 已证明是该 shape 类型的最优解)。 - ( - lambda m, n, k, aligned, **_kw: - aligned and (m * n > 2048 * 2048) and min(m, n) >= 64 - and ((m >= 16384 and max(n, k) <= 2048) - or (n >= 16384 and max(m, k) <= 2048)), - lambda: lambda: _hgemm_nn_kernel4[( - triton.cdiv(m, 128) * triton.cdiv(n, 256), - )]( - TensorDescriptor( - base=A, shape=[m, k], strides=[lda, 1], - block_shape=[128, 64], - ), - TensorDescriptor( - base=B, shape=[k, n], strides=[ldb, 1], - block_shape=[64, 256], - ), - TensorDescriptor( - base=C, shape=[m, n], strides=[ldc, 1], - block_shape=[128, 256], - ), - alpha, beta, m, n, k, beta_is_zero, - BLOCK_M=128, BLOCK_N=256, BLOCK_K=64, GROUP_M=8, - num_stages=4, num_warps=8, num_ctas=1, - ), - ), - - # ── 优先级 2 ───────────────────────────────────────────────── - # 对齐 + 大尺寸(m×n > 4M,min 维度 ≥ 64)。 - # 使用 kernel3,搭配 TensorDescriptor 和 autotuned configs - # (其 @libtuner 装饰器在运行时选择最优 BLOCK_M/N/K 等参数)。 - ( - lambda m, n, k, aligned, **_kw: - aligned and (m * n > 2048 * 2048) and min(m, n) >= 64, - lambda: lambda: _hgemm_nn_kernel3[grid]( - A, B, C, alpha, beta, m, n, k, - lda, ldb, ldc, beta_is_zero, +# ── 条件谓词(命名函数,非 lambda)─────────────────────────────── + +def _hgemm_nn_is_skinny_aligned_large(m, n, k, aligned, **_kw): + return ( + aligned + and (m * n > 2048 * 2048) + and min(m, n) >= 64 + and ( + (m >= 16384 and max(n, k) <= 2048) + or (n >= 16384 and max(m, k) <= 2048) + ) + ) + +def _hgemm_nn_is_aligned_large(m, n, k, aligned, **_kw): + return aligned and (m * n > 2048 * 2048) and min(m, n) >= 64 + +def _hgemm_nn_is_aligned_small(m, n, k, aligned, **_kw): + return aligned and max(m, n) <= 1024 + +def _hgemm_nn_is_default(**_kw): + return True + +# ── 工厂函数(接收 context 字典的键作为关键字参数)──────────────── + +def _hgemm_nn_build_kernel4(A, B, C, m, n, k, lda, ldb, ldc, + alpha, beta, beta_is_zero): + return lambda: _hgemm_nn_kernel4[( + triton.cdiv(m, 128) * triton.cdiv(n, 256), + )]( + TensorDescriptor( + base=A, shape=[m, k], strides=[lda, 1], + block_shape=[128, 64], ), - ), - - # ── 优先级 3 ───────────────────────────────────────────────── - # 对齐 + 小/中等尺寸(max ≤ 1024)。 - # 使用 kernel2,搭配 block_ptr 和 autotuned configs。 - ( - lambda m, n, k, aligned, **_kw: - aligned and max(m, n) <= 1024, - lambda: lambda: _hgemm_nn_kernel2[grid]( - A, B, C, alpha, beta, m, n, k, - lda, ldb, ldc, beta_is_zero, + TensorDescriptor( + base=B, shape=[k, n], strides=[ldb, 1], + block_shape=[64, 256], ), - ), - - # ── 优先级 4(兜底)────────────────────────────────────────── - # 其余所有情况:未对齐,或中等/大尺寸但未被上述条目覆盖。 - # 使用原始 level3 kernel,基于指针访问。 - ( - lambda **_kw: True, - lambda: lambda: _hgemm_nn_kernel[grid]( - A, B, C, alpha, beta, m, n, k, - lda, ldb, ldc, beta_is_zero, + TensorDescriptor( + base=C, shape=[m, n], strides=[ldc, 1], + block_shape=[128, 256], ), - ), + alpha, beta, m, n, k, beta_is_zero, + BLOCK_M=128, BLOCK_N=256, BLOCK_K=64, GROUP_M=8, + num_stages=4, num_warps=8, num_ctas=1, + ) + +def _hgemm_nn_build_kernel3(A, B, C, m, n, k, lda, ldb, ldc, + alpha, beta, beta_is_zero): + grid = lambda meta: ( + triton.cdiv(m, meta["BLOCK_M"]) * triton.cdiv(n, meta["BLOCK_N"]), + ) + return lambda: _hgemm_nn_kernel3[grid]( + A, B, C, alpha, beta, m, n, k, lda, ldb, ldc, beta_is_zero, + ) + +def _hgemm_nn_build_kernel2(A, B, C, m, n, k, lda, ldb, ldc, + alpha, beta, beta_is_zero): + grid = lambda meta: ( + triton.cdiv(m, meta["BLOCK_M"]) * triton.cdiv(n, meta["BLOCK_N"]), + ) + return lambda: _hgemm_nn_kernel2[grid]( + A, B, C, alpha, beta, m, n, k, lda, ldb, ldc, beta_is_zero, + ) + +def _hgemm_nn_build_kernel(A, B, C, m, n, k, lda, ldb, ldc, + alpha, beta, beta_is_zero): + grid = lambda meta: ( + triton.cdiv(m, meta["BLOCK_M"]) * triton.cdiv(n, meta["BLOCK_N"]), + ) + return lambda: _hgemm_nn_kernel[grid]( + A, B, C, alpha, beta, m, n, k, lda, ldb, ldc, beta_is_zero, + ) + +_HGEMM_NN_DISPATCH = StaticDispatch([ + # skinny + 对齐大尺寸 → kernel4(TensorDescriptor,硬编码 config) + (_hgemm_nn_is_skinny_aligned_large, _hgemm_nn_build_kernel4), + # 对齐大尺寸 → kernel3(TensorDescriptor,autotuned config) + (_hgemm_nn_is_aligned_large, _hgemm_nn_build_kernel3), + # 对齐小尺寸 → kernel2(block_ptr) + (_hgemm_nn_is_aligned_small, _hgemm_nn_build_kernel2), + # 默认 → kernel(原始实现) + (_hgemm_nn_is_default, _hgemm_nn_build_kernel), ]) +``` -runner = dispatch.lookup_and_build(m, n, k, aligned) -runner() +### 每次调用 —— 在 `hgemm()` 内部 + +```python +def hgemm(transa, transb, m, n, k, alpha, A, lda, B, ldb, beta, C, ldc): + # ... 参数校验 ... + beta_is_zero = beta == 0.0 + aligned = _is_gemm_aligned(A, lda, B, ldb, C, ldc) + + with torch_device_fn.device(A.device): + if transa == CUBLAS_OP_N and transb == CUBLAS_OP_N: + runner = _HGEMM_NN_DISPATCH.lookup_and_build( + m, n, k, aligned, + context=dict( + A=A, B=B, C=C, m=m, n=n, k=k, + lda=lda, ldb=ldb, ldc=ldc, + alpha=alpha, beta=beta, beta_is_zero=beta_is_zero, + ), + ) + runner() + # ... 其他 transa/transb 分支 ... ``` ### 调度逻辑总结 @@ -206,7 +259,7 @@ runner() 2. **顺序匹配**:条件自上而下求值;首次 `True` 即胜出。 3. **必须兜底**:最后一条必须匹配所有 shape(防止 `ValueError`)。 4. **条件互斥**:条目间不应重叠以保证行为可预测。高优先级条目应有更具体的条件。 -5. **双层 lambda 工厂**:搭配 `@libentry()` 装饰的 Triton kernel 时必须使用。内层 `lambda:` 将 `LibEntry.run()` 延迟到 `runner()` 时。 +5. **使用命名函数而非 lambda**:条件和工厂应为模块级别的命名函数,在 `StaticDispatch` 表中按名称引用。每次调用的变化数据(tensor、标量)通过 `context` 字典传入 `lookup_and_build()`,避免每次调用重新创建闭包。 --- @@ -241,37 +294,68 @@ runner() ### 常见模式 -**按维度分优先级**: +**按维度分优先级**(模块级命名函数): ```python -[ - (lambda m, *_kw, **__: m > 8192, factory_a), # 超大 - (lambda m, *_kw, **__: m > 1024, factory_b), # 大 - (lambda m, *_kw, **__: m > 256, factory_c), # 中等 - (lambda **_kw: True, factory_d), # 小 -] +def is_very_large(m, **_kw): + return m > 8192 + +def is_large(m, **_kw): + return m > 1024 + +def is_medium(m, **_kw): + return m > 256 + +def is_default(**_kw): + return True + +_DISPATCH = StaticDispatch([ + (is_very_large, build_kernel_a), + (is_large, build_kernel_b), + (is_medium, build_kernel_c), + (is_default, build_kernel_d), +]) ``` **按对齐分优先级**: ```python -[ - (lambda aligned, **_kw: aligned and is_large(**kw), aligned_large_factory), - (lambda aligned, **_kw: aligned, aligned_small_factory), - (lambda **_kw: True, unaligned_factory), -] +def is_aligned_large(aligned, m, n, k, **_kw): + return aligned and (m * n > 2048 * 2048) + +def is_aligned_only(aligned, **_kw): + return aligned + +def is_default(**_kw): + return True + +_DISPATCH = StaticDispatch([ + (is_aligned_large, build_aligned_large), + (is_aligned_only, build_aligned), + (is_default, build_fallback), +]) ``` **维度 + 对齐组合**(如 hgemm_nn): ```python -[ - (lambda aligned, m, n, k, **_kw: - aligned and meets_criteria_A(m, n, k), - factory_a), - (lambda aligned, m, n, k, **_kw: - aligned and meets_criteria_B(m, n, k), - factory_b), - (lambda **_kw: True, - fallback_factory), -] +def is_skinny_aligned_large(m, n, k, aligned, **_kw): + return (aligned and (m * n > 2048 * 2048) and min(m, n) >= 64 + and ((m >= 16384 and max(n, k) <= 2048) + or (n >= 16384 and max(m, k) <= 2048))) + +def is_aligned_large(m, n, k, aligned, **_kw): + return aligned and (m * n > 2048 * 2048) and min(m, n) >= 64 + +def is_aligned_small(m, n, k, aligned, **_kw): + return aligned and max(m, n) <= 1024 + +def is_default(**_kw): + return True + +_DISPATCH = StaticDispatch([ + (is_skinny_aligned_large, build_kernel4), + (is_aligned_large, build_kernel3), + (is_aligned_small, build_kernel2), + (is_default, build_kernel), +]) ``` --- @@ -308,6 +392,6 @@ runner() # 等价于 my_kernel_fn(A, B, C, alpha=1.0) 可以。例如,用 `StaticDispatch` 处理已知最优的 shape 类型,用 `SizeAutoDispatch` 处理其余情况。只需确保从正确的 dispatch 返回 runner 即可。 -### Q: 为什么 @libentry kernel 需要双层 lambda? +### Q: 为什么需要双层 lambda? -`@libentry()` 将 Triton `JITFunction` 包装为 `LibEntry` 对象。调用 `entry[grid](args...)` 会触发 `LibEntry.run()`,该方法编译(如果需要)、启动 kernel 并返回 `(kernel_obj, constexprs)`。内层 `lambda:` 将这个过程包装为可调用的 runner 而不立即执行 kernel。 +`@libentry()` 将 Triton `JITFunction` 包装为 `LibEntry` 对象。调用 `entry[grid](args...)` 会触发 `LibEntry.run()`,该方法编译(如果需要)、启动 kernel 并返回 `(kernel_obj, constexprs)`。返回的 `lambda:` 将这个过程包装为可调用的 runner 而不立即执行 kernel。 diff --git a/src/flag_blas/runtime/dispatch.py b/src/flag_blas/runtime/dispatch.py index 2d57a95e..b163340f 100644 --- a/src/flag_blas/runtime/dispatch.py +++ b/src/flag_blas/runtime/dispatch.py @@ -331,15 +331,43 @@ class StaticDispatch: Maps shape conditions directly to kernel factories. No autotune, no caching — the first matching entry wins immediately. - Usage:: + The dispatch table itself is designed to be created once at module + level and reused across calls. The per-call varying arguments (e.g. + tensors A, B, C) are passed to ``lookup_and_build`` via the + ``context`` dict. Factories receive these as keyword arguments, + avoiding the need to re-create closures on every call. + + Usage + ----- + Module level (once):: + + def is_aligned(m, n, k, aligned, **_kw): + return aligned + + def is_default(**_kw): + return True + + def build_aligned_runner(A, B, C, m, n, k, lda, ldb, ldc, + alpha, beta, beta_is_zero): + return lambda: _kernel2[...](A, B, C, alpha, beta, ...) - dispatch = StaticDispatch([ - (is_aligned, lambda: make_aligned_runner(A, lda, ...)), - (is_thin_large_m, lambda: make_thin_runner(A, lda, ...)), - (lambda **kw: True, lambda: make_fallback_runner(A, lda, ...)), + def build_fallback_runner(A, B, C, m, n, k, lda, ldb, ldc, + alpha, beta, beta_is_zero): + return lambda: _kernel[...](A, B, C, alpha, beta, ...) + + _MY_DISPATCH = StaticDispatch([ + (is_aligned, build_aligned_runner), + (is_default, build_fallback_runner), ]) - runner = dispatch.lookup_and_build(m, n, k, aligned) + Per-call (inside the function):: + + runner = _MY_DISPATCH.lookup_and_build( + m, n, k, aligned, + context=dict(A=A, B=B, C=C, m=m, n=n, k=k, + lda=lda, ldb=ldb, ldc=ldc, + alpha=alpha, beta=beta, beta_is_zero=beta_is_zero), + ) runner() Parameters @@ -348,12 +376,13 @@ class StaticDispatch: A list of ``(condition, factory)`` pairs evaluated **in order**. Each ``condition`` is a callable with signature ``(m, n, k, aligned, **extra) -> bool``. - Each ``factory`` is a zero-arg callable that returns a - ``Callable[[], None]`` runner. + Each ``factory`` is a callable that accepts keyword arguments + from ``context`` (or no arguments if ``context`` is None) and + returns a ``Callable[[], None]`` runner. The last entry should be a catch-all (condition always True). """ - _Entry = Tuple[Callable[..., bool], Callable[[], Callable[[], None]]] + _Entry = Tuple[Callable[..., bool], Callable[..., Callable[[], None]]] def __init__( self,