Merge pull request '请求合并' (#31) from xiao-ke/op_optimization:master into master

This commit is contained in:
Beckylu 2026-06-18 11:12:15 +08:00
commit 57f56aabb8
13 changed files with 908 additions and 331 deletions

View File

@ -0,0 +1,52 @@
---
sectionTitle: "接口约定"
type: "codeSample"
lang: "tilelang"
---
你必须在提交的 Python 代码中提供 `run_kernel` 函数,函数名、参数顺序、类型必须完全一致:
```python
import tilelang
import tilelang.language as T
from tilelang import jit
real_kernel = None
@jit
def build_kernel(*args):
@T.prim_func
def kernel(*args):
...
return kernel
def run_kernel(
q, # Tensor[bf16], shape (batch_size * seq_len, num_qo_heads, head_dim_qk)
k, # Tensor[bf16], shape (batch_size * seq_len, num_kv_heads, head_dim_qk)
v, # Tensor[bf16], shape (batch_size * seq_len, num_kv_heads, head_dim_vo)
output, # Tensor[bf16], shape (batch_size * seq_len, num_qo_heads, head_dim_vo)
qo_indptr, # Tensor[int32], shape (batch_size + 1)
kv_indptr, # Tensor[int32], shape (batch_size + 1)
batch_size, # int64
seq_len, # int64
num_qo_heads, # int64
num_kv_heads, # int64
head_dim_qk, # int64
head_dim_vo, # int64
causal, # int64
):
global real_kernel
if real_kernel is None:
real_kernel = build_kernel(...)
real_kernel(q, k, v, output, qo_indptr, kv_indptr,
batch_size, seq_len, num_qo_heads, num_kv_heads,
head_dim_qk, head_dim_vo, causal)
```
### 参数说明
* `q/k/v`FlashInfer ragged prefill 输入 tensor连续 `bfloat16`
* `output`:输出缓冲区,连续 `bfloat16`,需要写入结果
* `qo_indptr/kv_indptr`ragged indptr连续 `int32`
* `causal`:是否启用 causal mask评测中固定为 `1`
`run_kernel` 内部需要自行计算合适的 grid/block并 launch 你实现的 TileLang kernel。

View File

@ -0,0 +1,41 @@
---
sectionTitle: "接口约定"
type: "codeSample"
lang: "triton"
---
你必须在提交的 Python 代码中提供 `run_kernel` 函数,函数名、参数顺序、类型必须完全一致:
```python
import triton
import triton.language as tl
@triton.jit
def your_kernel(...):
...
def run_kernel(
q, # Tensor[bf16], shape (batch_size * seq_len, num_qo_heads, head_dim_qk)
k, # Tensor[bf16], shape (batch_size * seq_len, num_kv_heads, head_dim_qk)
v, # Tensor[bf16], shape (batch_size * seq_len, num_kv_heads, head_dim_vo)
output, # Tensor[bf16], shape (batch_size * seq_len, num_qo_heads, head_dim_vo)
qo_indptr, # Tensor[int32], shape (batch_size + 1)
kv_indptr, # Tensor[int32], shape (batch_size + 1)
batch_size, # int64
seq_len, # int64
num_qo_heads, # int64
num_kv_heads, # int64
head_dim_qk, # int64
head_dim_vo, # int64
causal, # int64
):
...
```
### 参数说明
* `q/k/v`FlashInfer ragged prefill 输入 tensor连续 `bfloat16`
* `output`:输出缓冲区,连续 `bfloat16`,需要写入结果
* `qo_indptr/kv_indptr`ragged indptr连续 `int32`
* `causal`:是否启用 causal mask评测中固定为 `1`
`run_kernel` 内部需要自行计算合适的 grid/block并 launch 你实现的 Triton kernel。

View File

@ -0,0 +1,55 @@
---
sectionTitle: "接口约定"
type: "codeSample"
lang: "tilelang"
---
你必须在提交的 Python 代码中提供 `run_kernel` 函数,函数名、参数顺序、类型必须完全一致:
```python
import tilelang
import tilelang.language as T
from tilelang import jit
real_kernel = None
@jit
def build_kernel(*args):
@T.prim_func
def kernel(*args):
...
return kernel
def run_kernel(
q, # Tensor[bf16], shape (batch_size * seq_len, num_qo_heads, head_dim)
kv_data, # Tensor[bf16], shape (num_blocks, 2, page_block_size, num_kv_heads, head_dim)
output, # Tensor[bf16], shape (batch_size * seq_len, num_qo_heads, head_dim)
qo_indptr, # Tensor[int32], shape (batch_size + 1)
kv_indptr, # Tensor[int32], shape (batch_size + 1)
kv_indices, # Tensor[int32], shape (num_blocks)
last_page_len, # Tensor[int32], shape (batch_size)
batch_size, # int64
seq_len, # int64
num_qo_heads, # int64
num_kv_heads, # int64
head_dim, # int64
page_block_size, # int64
causal, # int64
):
global real_kernel
if real_kernel is None:
real_kernel = build_kernel(...)
real_kernel(q, kv_data, output, qo_indptr, kv_indptr, kv_indices, last_page_len,
batch_size, seq_len, num_qo_heads, num_kv_heads,
head_dim, page_block_size, causal)
```
### 参数说明
* `q`query tensor连续 `bfloat16`
* `kv_data`paged KV cache连续 `bfloat16`
* `output`:输出缓冲区,连续 `bfloat16`,需要写入结果
* `qo_indptr/kv_indptr/kv_indices/last_page_len`paged KV metadata连续 `int32`
* `page_block_size`:评测中固定为 `16`
* `causal`:评测中固定为 `0`
`run_kernel` 内部需要自行计算合适的 grid/block并 launch 你实现的 TileLang kernel。

View File

@ -0,0 +1,44 @@
---
sectionTitle: "接口约定"
type: "codeSample"
lang: "triton"
---
你必须在提交的 Python 代码中提供 `run_kernel` 函数,函数名、参数顺序、类型必须完全一致:
```python
import triton
import triton.language as tl
@triton.jit
def your_kernel(...):
...
def run_kernel(
q, # Tensor[bf16], shape (batch_size * seq_len, num_qo_heads, head_dim)
kv_data, # Tensor[bf16], shape (num_blocks, 2, page_block_size, num_kv_heads, head_dim)
output, # Tensor[bf16], shape (batch_size * seq_len, num_qo_heads, head_dim)
qo_indptr, # Tensor[int32], shape (batch_size + 1)
kv_indptr, # Tensor[int32], shape (batch_size + 1)
kv_indices, # Tensor[int32], shape (num_blocks)
last_page_len, # Tensor[int32], shape (batch_size)
batch_size, # int64
seq_len, # int64
num_qo_heads, # int64
num_kv_heads, # int64
head_dim, # int64
page_block_size, # int64
causal, # int64
):
...
```
### 参数说明
* `q`query tensor连续 `bfloat16`
* `kv_data`paged KV cache连续 `bfloat16`
* `output`:输出缓冲区,连续 `bfloat16`,需要写入结果
* `qo_indptr/kv_indptr/kv_indices/last_page_len`paged KV metadata连续 `int32`
* `page_block_size`:评测中固定为 `16`
* `causal`:评测中固定为 `0`
`run_kernel` 内部需要自行计算合适的 grid/block并 launch 你实现的 Triton kernel。

View File

@ -0,0 +1,57 @@
---
sectionTitle: "接口约定"
type: "codeSample"
lang: "tilelang"
---
你必须在提交的 Python 代码中提供 `run_kernel` 函数,函数名、参数顺序、类型必须完全一致:
```python
import tilelang
import tilelang.language as T
from tilelang import jit
real_kernel = None
@jit
def build_kernel(*args):
@T.prim_func
def kernel(*args):
...
return kernel
def run_kernel(
q_nope, # Tensor[bf16], shape (batch_size, num_heads, head_dim_ckv)
q_pe, # Tensor[bf16], shape (batch_size, num_heads, head_dim_kpe)
ckv, # Tensor[bf16], shape (batch_size * seq_len, 1, head_dim_ckv)
kpe, # Tensor[bf16], shape (batch_size * seq_len, 1, head_dim_kpe)
output, # Tensor[bf16], shape (batch_size, num_heads, head_dim_ckv)
q_indptr, # Tensor[int32], shape (batch_size + 1)
kv_indptr, # Tensor[int32], shape (batch_size + 1)
kv_indices, # Tensor[int32], shape (batch_size * seq_len)
kv_lens, # Tensor[int32], shape (batch_size)
batch_size, # int64
seq_len, # int64
num_heads, # int64
head_dim_ckv, # int64
head_dim_kpe, # int64
page_size, # int64
causal, # int64
):
global real_kernel
if real_kernel is None:
real_kernel = build_kernel(...)
real_kernel(q_nope, q_pe, ckv, kpe, output,
q_indptr, kv_indptr, kv_indices, kv_lens,
batch_size, seq_len, num_heads,
head_dim_ckv, head_dim_kpe, page_size, causal)
```
### 参数说明
* `q_nope/q_pe/ckv/kpe`MLA attention 输入 tensor连续 `bfloat16`
* `output`:输出缓冲区,连续 `bfloat16`,需要写入结果
* `q_indptr/kv_indptr/kv_indices/kv_lens`paged attention metadata连续 `int32`
* `page_size`:评测中固定为 `1`
* `causal`:评测中固定为 `0`
`run_kernel` 内部需要自行计算合适的 grid/block并 launch 你实现的 TileLang kernel。

View File

@ -0,0 +1,45 @@
---
sectionTitle: "接口约定"
type: "codeSample"
lang: "triton"
---
你必须在提交的 Python 代码中提供 `run_kernel` 函数,函数名、参数顺序、类型必须完全一致:
```python
import triton
import triton.language as tl
@triton.jit
def your_kernel(...):
...
def run_kernel(
q_nope, # Tensor[bf16], shape (batch_size, num_heads, head_dim_ckv)
q_pe, # Tensor[bf16], shape (batch_size, num_heads, head_dim_kpe)
ckv, # Tensor[bf16], shape (batch_size * seq_len, 1, head_dim_ckv)
kpe, # Tensor[bf16], shape (batch_size * seq_len, 1, head_dim_kpe)
output, # Tensor[bf16], shape (batch_size, num_heads, head_dim_ckv)
q_indptr, # Tensor[int32], shape (batch_size + 1)
kv_indptr, # Tensor[int32], shape (batch_size + 1)
kv_indices, # Tensor[int32], shape (batch_size * seq_len)
kv_lens, # Tensor[int32], shape (batch_size)
batch_size, # int64
seq_len, # int64
num_heads, # int64
head_dim_ckv, # int64
head_dim_kpe, # int64
page_size, # int64
causal, # int64
):
...
```
### 参数说明
* `q_nope/q_pe/ckv/kpe`MLA attention 输入 tensor连续 `bfloat16`
* `output`:输出缓冲区,连续 `bfloat16`,需要写入结果
* `q_indptr/kv_indptr/kv_indices/kv_lens`paged attention metadata连续 `int32`
* `page_size`:评测中固定为 `1`
* `causal`:评测中固定为 `0`
`run_kernel` 内部需要自行计算合适的 grid/block并 launch 你实现的 Triton kernel。

View File

@ -0,0 +1,52 @@
---
sectionTitle: "接口约定"
type: "codeSample"
lang: "tilelang"
---
你必须在提交的 Python 代码中提供 `run_kernel` 函数,函数名、参数顺序、类型必须完全一致:
```python
import tilelang
import tilelang.language as T
from tilelang import jit
real_kernel = None
@jit
def build_kernel(*args):
@T.prim_func
def kernel(*args):
...
return kernel
def run_kernel(
q, # Tensor[bf16], shape (batch_size, num_qo_heads, head_dim)
kv_data, # Tensor[bf16], shape (num_blocks, 2, page_block_size, num_kv_heads, head_dim)
output, # Tensor[bf16], shape (batch_size, num_qo_heads, head_dim)
kv_indptr, # Tensor[int32], shape (batch_size + 1)
kv_indices, # Tensor[int32], shape (num_blocks)
last_page_len, # Tensor[int32], shape (batch_size)
batch_size, # int64
seq_len_kv, # int64
num_qo_heads, # int64
num_kv_heads, # int64
head_dim, # int64
page_block_size, # int64
):
global real_kernel
if real_kernel is None:
real_kernel = build_kernel(...)
real_kernel(q, kv_data, output, kv_indptr, kv_indices, last_page_len,
batch_size, seq_len_kv, num_qo_heads,
num_kv_heads, head_dim, page_block_size)
```
### 参数说明
* `q`decode query tensor连续 `bfloat16`
* `kv_data`paged KV cache连续 `bfloat16`
* `output`:输出缓冲区,连续 `bfloat16`,需要写入结果
* `kv_indptr/kv_indices/last_page_len`paged KV metadata连续 `int32`
* `page_block_size`:评测中固定为 `16`
`run_kernel` 内部需要自行计算合适的 grid/block并 launch 你实现的 TileLang kernel。

View File

@ -0,0 +1,41 @@
---
sectionTitle: "接口约定"
type: "codeSample"
lang: "triton"
---
你必须在提交的 Python 代码中提供 `run_kernel` 函数,函数名、参数顺序、类型必须完全一致:
```python
import triton
import triton.language as tl
@triton.jit
def your_kernel(...):
...
def run_kernel(
q, # Tensor[bf16], shape (batch_size, num_qo_heads, head_dim)
kv_data, # Tensor[bf16], shape (num_blocks, 2, page_block_size, num_kv_heads, head_dim)
output, # Tensor[bf16], shape (batch_size, num_qo_heads, head_dim)
kv_indptr, # Tensor[int32], shape (batch_size + 1)
kv_indices, # Tensor[int32], shape (num_blocks)
last_page_len, # Tensor[int32], shape (batch_size)
batch_size, # int64
seq_len_kv, # int64
num_qo_heads, # int64
num_kv_heads, # int64
head_dim, # int64
page_block_size, # int64
):
...
```
### 参数说明
* `q`decode query tensor连续 `bfloat16`
* `kv_data`paged KV cache连续 `bfloat16`
* `output`:输出缓冲区,连续 `bfloat16`,需要写入结果
* `kv_indptr/kv_indices/last_page_len`paged KV metadata连续 `int32`
* `page_block_size`:评测中固定为 `16`
`run_kernel` 内部需要自行计算合适的 grid/block并 launch 你实现的 Triton kernel。