跳到主要内容
起居室老虎
返回

AsyncT vllm适配、加速笔记(一)

3,293 字约 13 分钟

最近为了评估 AsyncT 的模型性能,需要跑各种越来越大越来越复杂的 Benchmark。之前是直接在 Megatron 里起一个 step=final 的训练脚本,用最后一次保存 ckpt 时候触发的 lm-eval 跑结果。但整个流程速度极慢,单独跑一个 gsm8k 就需要将近 40 分钟,效率太低了。加之最近正在思考这个显然更简单的新结构在 GPU 上能取得哪些现阶段就能够 claim 的优势,遂做此结构的推理框架适配和优化。

0. Preliminary

相比于 Softmax Transformer,AsyncT 主要的不同在于:

  1. 将 Transformer 中所有的 RMSNorm(pre-norm,qk-norm,final norm,以及下文提到的 Peri-LN 启发的 post-attn norm)替换为 DyHT:
DyHT(x)=y⊙clamp(ax,−1,1)\mathrm{DyHT}(x)=y\odot\mathrm{clamp}(ax, -1, 1)

是参考何恺明 Transformer without Normalization 中 DyT 的 SNN & Quantization 更友好的版本。他的论文中其实就提到过 DyT 有理论上比 RMSNorm 更高的计算效率(访存显著减少,并且 element-wise 操作更好 fuse),不过原论文中的实验做的比较 naive,知乎上也有人攻击11。

  1. 将 Softmax Attention 替换为:
𝐒h=1d𝐐h𝐊h⊤⊙𝐌casualAsyncI(𝐒;σ)=clamp(𝐒−σ⋅CasualAvgPooling1D(ReLU(𝐒),0,1)𝐀h=AsyncI(𝐒h;σh)𝐎h=𝐀h𝐕h\begin{align*} \mathbf{S}_h&=\frac{1}{\sqrt{d}}\mathbf{Q}_h\mathbf{K}_h^\top\odot \mathbf{M}_\text{casual}\\ \mathrm{AsyncI}(\mathbf{S};\sigma)&=\mathrm{clamp}(\mathbf{S}-\sigma\cdot\mathrm{CasualAvgPooling1D}(\mathrm{ReLU}(\mathbf{S}), 0, 1)\\ \mathbf{A}_h&=\mathrm{AsyncI}(\mathbf{S}_h;\sigma_h)\\ \mathbf{O}_h&=\mathbf{A}_h\mathbf{V}_h \end{align*}

相比于 Softmax,主要区别在于(1)将其中的 Softmax 算子移除替换为了 AsyncI,进一步移除了模型中的 normalization;(2)由于移除 normalization 后 Attention Score 之和为 1 的约束一同消失,将 Softmax 中原本为dh\sqrt{d_\text{h}}的 scaling factor 修改为d\sqrt{d}。

  1. 在 Attention 之后、Residual 之前添加了一个额外的 DyHT,用于稳定模型输出。此方法的启发来源于 Peri-LN,没有像原文一样在 MLP 之后也添加是因为测试出来效果不算很好。
  2. 将 FFN 中的 SwiGLU 替换为 ReGLU,主要区别是将其中的 SiLU 替换为了 ReLU。一开始的目的是为了避免构造差分形式的 SNN 算子,不过后面 ReLU Strikes Back 和 Sparsing Law 等论文指出,使用 ReGLU 在常见的 Benchmark 上并不会引入特别多的 perf degradation,同时还有稀疏度更高、硬件更友好(不用 SFU 了)等优势性质。

通过以上算子的修改,AsyncT 相比于 Softmax(可能,或者说我们希望)具有以下优点:

  1. LoCC Friendly,算子完全适配我们在 ViStream 中提出的 LoCC Law,移除 normalization 之类的 synchronous 算子之后模型能够非常流畅地转换为能够在异步硬件上部署的 SNN;
  2. Quantization Friendly,之前分享的 A Unified View of Attention and Residual Sinks: Outlier-Driven Rescaling is Essential for Transformer Training 指出,LLM 中的 outlier 往往是由于模型尝试利用 softmax/rmsnorm 中的 normalization 进行 rescaling 产生的,而我们不仅移除了 normalization,还通过 DyHT 强硬地限制了每个元素位置上输出的数值范围,保证了模型的稳定性,理想情况下量化也会非常容易;
  3. Sparse Friendly,这个 friendly 的叫法有点强行,不过将 SwiGLU 替换为 ReGLU 确实观察到了模型的稀疏性显著提高的情况,下面这个模型在训练中就已经发现了 MLP 中~70% 的平均稀疏度,并且按 Sparsing Law 的描述模型继续训练会变得更加稀疏。

image.png

可以看到总体而言 AsyncT 相比于 Softmax Model 的差别还是非常大的,比较意外地是这个模型在我们的 setup 下观察到了可以说是接近 Softmax Attention Model 的能力表现,还是比较惊喜的。那接下来的重点就是如何将上面这个模型嵌入到 vllm 中,vllm 中这么多成熟的 technique 有多少能在 AsyncT 上直接应用就能得到收益,而有哪些是为了 softmax model 设计的在 AsyncT 上需要修改、反而可以利用 AsyncT 自己的性质来减少麻烦等等。

1. VLLM 适配

1.1. PyTorch 基本版

要做适配当然要从一个最基本的 PyTorch 版本开始,起码测一测保证里面该换的东西都能换,之后也有一个精度 Baseline 和一个优化的对比。好在 VLLM 已经是个成熟的大框架,支持的各种 Backend 本来就非常多,要做这个适配基本上核心工作量都在写 AttentionMetadataBuilder(每步把 block table、seq lens 等元信息打包)和 AttentionImpl.forward(拿到 Q/K/V 和 metadata,写出 attention output),其他部分还是些小修小改。

具体而言,我们需要实现的内容包括:

KV cache 的 layout 先直接对齐 FlashAttention:[2, num_blocks, block_size, num_kv_heads, head_size](第 0 维是 K/V)。这个选择有两个理由:

  1. reshape_and_cache_flash 这个 op 已经在 _C 里编译好了,写 KV 直接拿来用就行;
  2. 后面想偷 FA 的某些工具(比如 flash_attn varlen prefill 的 cache 路径)会更容易。

代价是访存模式不如 vLLM 自家的 paged_attention_v2 那种 head-major + interleaved-x(差大概 10-20%),不过这只是作为 baseline 的一个基本实现,所以无关紧要。

for req in requests:
gather K/V
build mask
hp1_attention_eager()
scatter output

这个纯 PyTorch 实现跑 gsm8k 要 2131 秒,HF eager 只要 771 秒——vLLM 比 HF 慢了 2.76 倍。vLLM 的强项是把大量 request 的 decode/prefill 合并成连续 GPU workload;而这个版本等于把 vLLM 的 batch 又拆回了 Python request loop。profile 结果指出:vLLM 每个 forward step 会有约 250 个并发 request × 35 层 attention = 约 9000 次 Python 迭代。AttentionImpl.forward 本身就在 Python 端做 per-request 的 block_table[req_idx, :num_blocks]、k_flat[tok_indices]、调用 hp1_attention_eager_kv(...)、output[q_start:q_end] = ...这一连串操作。每一步都在 GPU 上 launch 小 kernel 然后等 host 侧 Python 决策,整个 pipeline 完全是 launch-overhead-bound 的。

所以后面这一轮优化的主线就变成:把 request 维度从 Python loop 里拿出来,变成 batched tensor / GPU kernel / graph replay 的一部分。

1.2. Batched Padded Matmul

既然 per-request Python loop 是 bottleneck,就把所有 request 合并成一个 padded batch 做一次性的 matmul:

# 先扫一遍算出最大长度
for req_idx in range(num_reqs):
max_t_q = max(max_t_q, t_q)
max_t_kv = max(max_t_kv, t_kv)
# 分配padded tensor
Q_padded = torch.zeros(B, max_t_q, H, D, ...)
K_padded = torch.zeros(B, max_t_kv, Hkv, D, ...)
# 用memcpy-only的Python loop填充(无per-request matmul)
for b, (req_idx, q_start, t_q, t_kv) in enumerate(req_info):
Q_padded[b, :t_q] = query[q_start:q_start+t_q]
# 关键:tok_indices一次构造,然后k_flat[tok_indices]做gather
K_padded[b, :t_kv] = k_flat[tok_indices]
# 一次batched call
Out_padded = asynct_attention_eager(Q_padded, K_padded, V_padded, ...)

注意这里有个微妙之处:Python loop 还在,但 loop 体里没有任何 matmul。loop 只做 Q_padded[b, :t_q] = query[q_start:q_start+t_q]这种 memcpy。matmul 被推到了 hp1_attention_eager(...)这一次 batched call。

通过上面这个简单的 batch 做法,arc_easy 加速了 2.21 倍,但 gsm8k OOM 了——256 个请求 × 2048 max_t_kv× 4 个 KV head × 128 维 × 2 字节 × 35 层 = 数 GB 级别的 padded K/V tensor。所以接下来要做长度分桶。

1.3. Length Bucketed Batching

核心 idea 是按 t_kv 排序,然后 greedy 地把 request 塞进 chunk,确保每个 chunk 的 B_chunk * chunk_max_t_kv 不超过预算(131072 tokens):

sorted_by_kv = sorted(range(len(req_info)), key=lambda i: req_info[i][3])
for i in sorted_by_kv:
new_max = max(cur_max_tkv, t_kv_i)
if (len(cur) + 1) * new_max > BUDGET_TOKENS:
chunks.append(cur)
cur = [i]

这样改完还有个好处:短 request 聚在一起时,它们所在 chunk 的 max_t_kv 也很小,padding 浪费骤减。gsm8k 从 2131s 降到 928s,同时也修了 OOM。HF eager 的差距从 2.76× 缩到 1.20×。

但这里仍然有两个问题:

  1. 每个 chunk 内仍然要准备 padded tensor;
  2. 每个 request 的 K/V gather 仍然是在 Python loop 中一个一个做。

1.4. Vectorized Gather

可以看到上一轮代码里 loop 中还有一些 index-only 的操作,那可以把它们也 vectorize。把 block table 展开成 [B_c, chunk_max_t_kv]的 token index 矩阵:

tok_indices =
block_table_chunk.view(B_c, max_blocks, 1) * block_size
+ offsets.view(1, 1, block_size)

然后一次性:

K_gathered = k_flat[tok_indices]
V_gathered = v_flat[tok_indices]

Q 也用类似方式 gather。这样 Python loop 只剩少量 chunk 级逻辑,不再每个 request 发起独立 gather/matmul。结果 gsm8k 到 457s,比 HF eager 快 1.69x。

1.5. Python 侧一些其他的尝试

前面提到 decode indice 是可以通过一次 build 给所有层复用的,但当前的代码每次做 attention 的时候都重复构建

q_decode_rows
seq_lens_dec
block_table_dec
kv_indptr
kv_indices
mid_o
filter_w

等 decode indices,因为 vLLM 一次 forward 中所有层共享同一个 attn_metadata,这些完全可以只构建一次,挂到 metadata 上。

之后又尝试启动 CUDA Graph,结果反而整体反而变慢。可以看出在 Kernel 执行时长还是显著的瓶颈的时候,启动 CUDA Graph 不一定是个有意义的操作。之后在做 CUDA Graph 等的修改之前,一定要先 profile、确认各种 kernel launch 的间隔相比于 kernel 本身执行的时间已经不可忽略了,再考虑这样的优化。

2. Triton Decode Kernel

PyTorch 端很难融合 FIR + clamp + p@V 这串 element-wise + reduce 的混合 op,开始快速把之前 Megatron 里的 Triton Kernel 搬进来,先从 Decode 开始。

2.1. Naive Two Stage Decode Triton Kernel

整个 Triton path 是两阶段 kernel 设计:

Stage 1:每个(req, q_head, split)组合一个 program,在该 split 的 K 范围内:

  1. 加载 Q [HEAD_SIZE]
  2. 构造 banded-Toeplitz FIR matrix [BLOCK_N, BLOCK_N]——这里有个 trick,HP1 的因果 FIR 可以用 Toeplitz 矩阵的 tl.dot 表达
  3. 按 BLOCK_N tile 扫 K:s = q · K^T → relu → tl.dot(·, FIR_toeplitz) → lp_scale → s_hp → clamp gate → p · V
  4. 写到 Mid_O_ptr[req, q_head, split_id, head_size]

Stage 2:sum reduce 各个 split 的 partial output。因为 p(HP1 的 gate 输出)和 V 的乘积是线性的,各 split 的 partial output 直接相加即可,不需要像 softmax 那样追加 exp scale matching。

@triton.jit
def _hp1_decode_stage1_kernel(
Q_ptr, # FULL query buffer [num_q_tokens, H, D]
Q_Rows_ptr, # [num_decode_reqs] int64 row in query for each decode req
...
):
cur_req = tl.program_id(0)
q_row = tl.load(Q_Rows_ptr + cur_req) # indirect indexing
q = tl.load(Q_ptr + q_row * stride_q_t + ...).to(tl.float32) * qk_scale
...
@triton.jit
def _hp1_decode_stage2_kernel(
Mid_O_ptr,
Out_ptr, # FULL output buffer
Q_Rows_ptr, # row in Out_ptr to write to
...
):
cur_req = tl.program_id(0)
out_row = tl.load(Q_Rows_ptr + cur_req)
# ... sum across splits, then store to Out_ptr[out_row, head, :]

这里有个相对 Triton 原生不友好的部分:

lowpass[k] = avg(relu(score[k]), relu(score[k-1]), relu(score[k-2]))

在 Triton 里最直接的写法是构造一个 banded Toeplitz 矩阵

𝐅avg=13[11100⋯00001110⋯00000111⋯000⋮⋱⋱⋱⋮00000⋯11100000⋯01100000⋯001]N×N\mathbf{F}_{\mathrm{avg}} = \frac{1}{3}\, \begin{bmatrix} 1 & 1 & 1 & 0 & 0 & \cdots & 0 & 0 & 0 \\ 0 & 1 & 1 & 1 & 0 & \cdots & 0 & 0 & 0 \\ 0 & 0 & 1 & 1 & 1 & \cdots & 0 & 0 & 0 \\ \vdots & & & \ddots & \ddots & \ddots & & & \vdots \\ 0 & 0 & 0 & 0 & 0 & \cdots & 1 & 1 & 1 \\ 0 & 0 & 0 & 0 & 0 & \cdots & 0 & 1 & 1 \\ 0 & 0 & 0 & 0 & 0 & \cdots & 0 & 0 & 1 \end{bmatrix}_{N\times N}

然后用 tl.dot:

relu_scores [BLOCK_N] × Toeplitz[FIR_K] -> lowpass[BLOCK_N]

这样做代码逻辑最简单,边界控制也方便。对于 Kernel Size = 3 的情况来说,上面的矩阵显然大部分位置都是 0,用一个矩阵乘来实现 shift-add 的效率实在是太低了。然而,Triton 的寄存器 tensor 模型不太方便表达 “从寄存器 tile 里按 k-i shift 取值” 的直接 FIR,尤其还要处理 split 边界的 FIR_K-1left overlap。

2.2. 带宽优化

做简单 profile,发现 K/V Load 是带宽的主要任务。两个比较直觉性的想法:

  1. 将所有的 Q Head 一起打包成 acc[H_q, HEAD_SIZE],在一个 triton kernel 执行过程中一次性做 reduce,节省 KV Cache 的读取。结果效率反而大幅度降低;经检查是打包后 register 数量大幅度增加,导致反而溢出到 local memory 上;
  2. 只在 GQA Shared 的 head 上共享 KV Cache Load,但现在的模型结构只有 GQA2,K/V bandwidth saving = T_kv × Hkv × D × 2 bytes / 2 ≈ 14μs/layer at H200 4.8 TB/s peak,好像没啥用,遂放弃。

3. Naive CUDA Kernel

Triton Kernel 的各种写法在之前 MegatronLM 上写训练算子的时候研究的已经很多了,限制还是太多了:

再在 Triton Level 上研究各种尝试对现在这个 Kernel 的形态意义不大。

3.1. First CUDA Kernel

复用 vLLM 的 paged_attention_v2_kernel 模板:

__shared__ float logits[MAX_SEQ_LEN + FIR_K - 1]; // 前面留Kenrel Size - 1个history slot
__shared__ float p_buf[MAX_SEQ_LEN];
for (int t = tid; t < seq_len; t += NUM_THREADS) {
const int idx = t + FIR_K - 1;
float fir_sum = 0.0f;
#pragma unroll
for (int k = 0; k < FIR_K; k++) {
float v = logits[idx - k]; // ← 直接在shared mem里下标访问,方便多了
if (hp_relu_pre) v = fmaxf(v, 0.0f);
fir_sum += v;
}
// ... gate compute ...
p_buf[t] = gamma_v * p_val;
}

跑 BS=32, T_kv=1024 的 Micro Bench 相比于 Triton 版本立刻获得~20% 的提速,效果明显。但 gsm8k 实际上下降 6%(385s vs 362s)。原因:前面的 dispatch 只在 num_prefills == 0 时启动 CUDA path(即 pure decode),mixed batch 全部 fallback 到 Triton;而 gsm8k 因为 continuous batching 大量是 mixed batch。

3.2. Mixed-batch Dispatch

把 dispatch CUDA Kernel 的条件,从 num_prefills == 0 改成 num_decodes > 0。CUDA path 处理 mixed batch 里的 decode subset,Path B Triton/eager 只处理 prefill subset。

if (_cuda_enabled and attn_metadata.num_decodes > 0 and H * D == 1024):
num_decodes = int(attn_metadata.num_decodes)
torch.ops._C.hp1_decode_paged(
output[:num_decodes], query[:num_decodes],
key_cache, value_cache,
attn_metadata.block_table[:num_decodes].to(torch.int32),
...
)
if attn_metadata.num_prefills == 0:
return output # pure decode — fully done
# fall through to Path B for prefill subset

vLLM 的 scheduler 把 decode request 和 prefill request 在同一个 batch 里排在一起(decode 在前,prefill 在后),所以 output[:num_decodes]和 query[:num_decodes]就刚好切到 decode subset。

3.3. Kernel Fusion

既然已经开始写 CUDA Kernel 了,就给其他的如 UClip、ReGLU 等都写一下实现、控制一下 kernel launch 的情况:

template <typename scalar_t, int NUM_THREADS,
bool ALPHA_PER_CHANNEL, bool GAMMA_PER_CHANNEL>
__global__ void uclip_norm_kernel(...) {
for (int d = tid; d < hidden_size; d += NUM_THREADS) {
float xf = uclip::to_float(x_row[d]);
float a = ALPHA_PER_CHANNEL ? alpha[d] : alpha[0];
float g = GAMMA_PER_CHANNEL ? gamma[d] : gamma[0];
float v = fminf(fmaxf(a * xf, clip_min), clip_max) * g;
y_row[d] = uclip::from_float<scalar_t>(v);
}
}

不过要注意:

最后的效果其实和 PyTorch Eager 版本差距也没有那么大。

ReGLU 同理,用一个 relu_and_mul 替代 relu(gate) * up 两步。

做了 Kernel Fusion 整体的 inference 能力已经大幅度提升,此时端到端的性能已经基本赶上 Flash Attention2 的速度。

3.4. CUDA Graph Again

这时候很自然地想,既然之前是 Kernel 性能受限,现在 Kernel 性能应该没什么问题了,总可以用了?结果启动 CUDA Graph 后捕获成功,但总吞吐继续变慢,还没有仔细做 Benchmarking,看来这边还有一些比较神秘的问题22。

至此 Part 1 的故事到一个阶段性节点:从 per-request Python loop = 2131s 走到 CUDA decode + Triton fallback = 277,7.7× wall time 减少,纯靠对 vLLM 架构(AttentionMetadataBuilder timing、KV cache layout、layer-level dispatch)的细致挖掘。

到这里,基本上算是实现了我们第一阶段的目标,此算子能够非常轻松地支持我们对各种 Benchmark 评估过程的支持、再也不用像之前一样等一两个小时跑一个小模型的 eval 结果、还需要大量占用卡的显存了。不过在优化的过程中我们实际上能够发现,AsyncT 在算子优化领域上实际上具有非常大的优化潜力,这和它的无 normalization(or async)的数学形态是息息相关的。下一节进入 CUDA 上的算子优化、Hopper 架构特性的完整利用。

脚注

  1. 如何评价 Meta 新论文 Transformers without Normalization? - 寒月灼华的回答 - 知乎 https://www.zhihu.com/question/14925347536/answer/124311637540 ↩

  2. 不过后面会发现,这实际上是 Benchmarking 存在问题,其中的 warmup 做法不太干净,此处的 graph mode 实际上已经可以引入不少收益了。 ↩