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 }} diff --git a/docs/static_dispatch.md b/docs/static_dispatch.md new file mode 100644 index 00000000..31d55450 --- /dev/null +++ b/docs/static_dispatch.md @@ -0,0 +1,397 @@ +# 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. + +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 + +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 + +Signature: `(context_key1=..., context_key2=..., ...) -> Callable[[], None]` + +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()`. + +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`. + +**Critical**: When used with `@libentry()`-decorated Triton kernels, the factory must return the kernel call wrapped in a **lambda**: + +```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 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. + +--- + +## Architecture + +Unlike `SizeAutoDispatch` with its multi-tier cache, `StaticDispatch` has a minimal design: + +``` +lookup_and_build(m, n, k, aligned, *, context, **extra) + │ + ├─ 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) + +--- + +## 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, *, context=None, **extra) +runner() +``` + +| Parameter | Type | Description | +|-----------|------|-------------| +| `m` | `int` | M dimension | +| `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. + +**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. + +### Module level — defined once + +```python +from flag_blas.runtime.dispatch import StaticDispatch +from triton.tools.tensor_descriptor import TensorDescriptor + +# ── 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], + ), + 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), +]) +``` + +### 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 + +| 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. **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. + +--- + +## 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** (module-level named functions): +```python +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 +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 +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), +]) +``` + +--- + +## 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 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 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 new file mode 100644 index 00000000..cfccfeab --- /dev/null +++ b/docs/static_dispatch_cn.md @@ -0,0 +1,397 @@ +# StaticDispatch 静态调度表 + +[源码: `src/flag_blas/runtime/dispatch.py`](../src/flag_blas/runtime/dispatch.py) + +`StaticDispatch` 是开发者维护的静态 kernel 调度表,将 shape 条件直接映射到预定的 kernel 工厂函数。**不进行 autotune、不 benchmark、不缓存** —— 条件按顺序求值,首次命中即返回。 + +调度表设计为模块级别创建一次,跨调用复用。每次调用的变化数据(tensor、标量等)通过 `context` 字典传入 `lookup_and_build()`,工厂函数无需通过闭包捕获变量。 + +--- + +## 适用场景 + +当**最优的 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(工厂函数) + +签名:`(context_key1=..., context_key2=..., ...) -> Callable[[], None]` + +一个命名函数(推荐使用命名函数而非 lambda),通过关键字参数接收每次调用的变化数据,返回一个 **runner** —— 执行 kernel 的零参数可调用对象。参数通过 `lookup_and_build()` 中的 `context` 字典传入。 + +当调度表在模块级别创建时,工厂函数只是对命名函数的纯引用 —— 没有闭包,没有每次调用重新构建的 lambda。每次调用的数据通过 `context` 流入。 + +**关键要点**:当搭配 `@libentry()` 装饰的 Triton kernel 使用时,工厂必须返回一个被 **lambda** 包装的 kernel 调用: + +```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_obj, constexprs)` 元组,而非可调用的 runner。`lambda:` 将执行延迟到 `runner()` 调用时。 + +--- + +## 架构 + +与 `SizeAutoDispatch` 的多级缓存架构不同,`StaticDispatch` 设计极简: + +``` +lookup_and_build(m, n, k, aligned, *, context, **extra) + │ + ├─ 条目 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`(最后一条应为兜底条目) + +--- + +## 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, *, context=None, **extra) +runner() +``` + +| 参数 | 类型 | 说明 | +|------|------|------| +| `m` | `int` | M 维度 | +| `n` | `int` | N 维度 | +| `k` | `int` | K 维度 | +| `aligned` | `bool` | 输入是否内存对齐 | +| `context` | `dict` 或 `None` | 每次调用的变化数据(tensor、标量等),以关键字参数形式传入匹配的工厂函数。为 `None` 时工厂函数无参调用。 | +| `**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 + +# ── 条件谓词(命名函数,非 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], + ), + 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 + 对齐大尺寸 → 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), +]) +``` + +### 每次调用 —— 在 `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 分支 ... +``` + +### 调度逻辑总结 + +| 优先级 | 条件 | 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**:条件和工厂应为模块级别的命名函数,在 `StaticDispatch` 表中按名称引用。每次调用的变化数据(tensor、标量)通过 `context` 字典传入 `lookup_and_build()`,避免每次调用重新创建闭包。 + +--- + +## 与 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 +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 +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 +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), +]) +``` + +--- + +## 辅助类: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: 为什么需要双层 lambda? + +`@libentry()` 将 Triton `JITFunction` 包装为 `LibEntry` 对象。调用 `entry[grid](args...)` 会触发 `LibEntry.run()`,该方法编译(如果需要)、启动 kernel 并返回 `(kernel_obj, constexprs)`。返回的 `lambda:` 将这个过程包装为可调用的 runner 而不立即执行 kernel。 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..8230d7b0 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 @@ -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, @@ -1859,69 +1949,17 @@ 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 + 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, + ), ) - 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 - ) + 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..b163340f 100644 --- a/src/flag_blas/runtime/dispatch.py +++ b/src/flag_blas/runtime/dispatch.py @@ -324,6 +324,93 @@ 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. + + 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, ...) + + 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), + ]) + + 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 + ---------- + 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 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]]] + + def __init__( + self, + table: List[_Entry], + ): + self._table = table + + def lookup_and_build( + self, + m: int, + 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 " + f"m={m}, n={n}, k={k}, aligned={aligned}" + ) + + class KernelRunner: """ A callable wrapper that executes a kernel with pre-bound arguments. 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