本仓库实现了一个基于 Transformer + Sparse Mixture-of-Experts (SMoE) 的点击率预测模型,针对 A800 推理场景做了深度优化。
Raw Batch (user sequences)
↓
RepEncoder: 28 slot embedding → LayerNorm → Linear(14336→512)
↓
TransformerEncoder: 8 × ( FlashAttn + SMoE ) layers
↓
Prediction Head: Linear(512→1) → sigmoid → click probability
| 键 | 类型 | Shape | 说明 |
|---|---|---|---|
user_offsets |
LongTensor | [num_users + 1] |
用户序列边界,[0, len1, len1+len2, ..., total_tokens] |
1 ~ 28 |
(values, offsets) |
values: [sum_of_feats]offsets: [total_tokens + 1] |
28 个特征槽,每个槽用 EmbeddingBag 的稀疏格式 |
pred_mask |
BoolTensor | [total_tokens] |
标记哪些 token 需要预测 |
logid |
LongTensor | [total_tokens] |
样本 ID,用于输出对齐 |
关键点:多个用户的序列首尾相接打包成一条长序列(total_tokens 个 token),每个用户内部保持时序因果。
输入:batch 字典
输出:seq_input → [total_tokens, 512]
操作流程:
-
Batched Embedding Bag(Triton 融合算子):
- 对 28 个槽位并行调用
triton_batched_embedding_bag,一次 kernel launch 完成全部查表。 - 每个槽位内的多个特征 ID 做 sum pooling(
mode="sum")。 - 输出
[batch_size, 28, 512]→ reshape 为[batch_size, 14336]。
- 对 28 个槽位并行调用
-
LayerNorm + Linear:
fused_embs = triton_batched_embedding_bag(...) # [batch_size, 28*512] normed = LayerNorm(fused_embs) # [batch_size, 14336] seq_input = Linear(normed) # [batch_size, 512]
优化:
- 零拷贝指针传递:只传 28 个
.data_ptr()给 Triton,避免torch.cat拷贝开销。 - 64 位寻址:vocab_size=5M × emb_dim=512 超 int32,用
tl.int64防溢出。 - 性能提升:查表时间从 0.918ms → 0.102ms(9× 加速)。
输入:seq_input → [total_tokens, 512],user_offsets
输出:[total_tokens, 512]
每层结构(共 8 层):
for layer in range(8):
# 1. Multi-Head Self-Attention (因果 + 跨用户隔离)
residual = x
x = LayerNorm(x)
qkv = Linear_qkv(x) # [total_tokens, 1536]
attn_out = FlashAttention_Varlen(qkv, user_offsets) # Triton 变长因果注意力
x = residual + Linear_out(attn_out)
# 2. Sparse Mixture-of-Experts (Top-2 路由,8 专家)
residual = x
x = LayerNorm(x)
topk_idx, topk_weight = Gate(x) # CUDA kernel: softmax→top-k→renorm
moe_out = FusedMoE(x, topk_idx, topk_weight) # Triton 融合算子
x = residual + moe_out为什么需要:
- 打包序列是
B=1, S=total_tokens的长序列,user_offsets标记每个用户段的边界。 - 标准 SDPA 需要物化
[S, S]的块对角因果掩码,44k tokens 时需 14.6 GiB 显存 → OOM。 - 外部
flash-attn需要预编译 wheel,部署脆弱。
算子设计:
- 并行网格:
(cdiv(max_seqlen, 64), num_users, n_heads)→ 每个 Program 处理一个用户段的一个 query 块。 - 段边界寻址:kernel 内从
cu_seqlens读seq_start / seq_end,直接定位该用户的 Q/K/V,无需 padding。 - 因果剪枝:query 第
i行只看 key0..i,K/V 循环在对角块处提前停止,省掉上三角计算。 - Online Softmax:维护
m_i(行最大值)、l_i(归一化分母)、acc(输出累加器),在 fp32 中逐块更新,最后写回 bf16。
性能收益(6823 tokens / 60 users):
- SDPA(含掩码构建):3.472 ms
- Triton FlashAttn:0.111 ms(31× 加速)
结构:
- Gate:
Linear(512 → 8)→ CUDA kernel 一次完成softmax + top-2 + renormalize。 - 8 个 Expert:每个 Expert 是
Linear(512 → 1024) + ReLU + Linear(1024 → 512)。 - 路由策略:每个 token 选 2 个专家,按 softmax 权重加权聚合。
Triton 融合优化(fused_moe_forward):
-
路由准备(GPU 异步):
sorted_route_ids, expert_ids, num_tokens_post_pad = prepare_routing_triton(...) # - 用 scatter_add 统计每个专家分到的 token 数 # - CUDA kernel 一次完成:专家 ID 填充 + token 重排 + route_positions 记录
-
W1 矩阵乘 + ReLU 融合:
hidden = grouped_matmul_triton(x, w1, ..., activation="relu") # Triton kernel 内: # - gather 输入 token(通过 sorted_route_ids) # - 按专家分组 GEMM (w1) # - ReLU 激活直接融入 kernel,避免单独 launch # - Early-Exit:grid 超出实际 token 数的 block 直接 return
-
W2 矩阵乘 + scatter 输出:
route_out = grouped_matmul_scatter_routes_triton(hidden, w2, ...) # 直接写入 route-major layout:row `token_id * top_k + k_idx`
-
加权归约:
out = route_weighted_sum_triton(route_out, topk_weights, ...) # 每个 token 对其 top-2 路由的输出做加权和
关键优化:
- 编译期常量:
N=1024、K=512声明为tl.constexpr,Triton 编译器展开循环边界,提升吞吐。 - Workspace 预分配:每层维护独立 workspace,复用
grouped_token_ids、hidden、route_out等 buffer,减少多层推理时的显存碎片。 - 动态配置:小 batch 用
BLOCK_M=16,大 batch 用BLOCK_M=64,最小化 padding 开销。
输入:encoder_output → [total_tokens, 512]
输出:logits → [total_tokens, 1]
logits = Linear(encoder_output) # [total_tokens, 1]
logits = torch.clamp(logits, -15.0, 15.0) # 防止数值爆炸
probs = torch.sigmoid(logits) # 点击概率最后用 pred_mask 提取需要预测的 token,配对 logid 输出。
瓶颈:每个 batch 约 211 次 kernel launch(8 层 × attention + MoE 路由的碎 kernel),CPU dispatch 开销 ~91ms/batch,而 GPU 实际计算只占 ~9ms/batch → launch-bound。
CUDA Graph 把整段计算图捕获一次,replay 时折叠成单次 dispatch,直接抹掉 launch 开销。
难点:CUDA Graph 要求静态 shape + 静态地址,但打包后的 token 数在 1221~10793 之间变化。
解决:
-
分桶 padding:以 512 为步长设 10 个桶(2560, 3072, ..., 7168),覆盖 98.6% batch。
- 每个桶惰性捕获一张独立的图,作用在固定
[S_PAD, 512]的静态 buffer 上。 - 超过顶桶(1.4% 长尾)回退 eager,避免极少数大 batch 撑爆显存。
- 每个桶惰性捕获一张独立的图,作用在固定
-
固定 attention grid:
- 捕获时用固定
max_seqlen = S_PAD决定 FlashAttn 的 grid 维度。 - Triton kernel 的逐块 early-exit(
q_block_start >= seqlen直接 return)让"偏大的 grid"对真实较短的段成为空操作,结果完全不变。
- 捕获时用固定
-
去除 host-sync:
- 原
max_seqlen = int(lengths.max().item())的.item()会强制同步,无法进图。 - 分桶后每桶
max_seqlen是常量,直接传入。
- 原
-
变长 num_seqs:
cu_seqlens长度固定为num_seqs+1=21。- 偶发的短 batch(<20 用户)用末尾 offset 在 GPU 上填充成零长段,被 kernel early-exit 跳过,不引入
.item()同步。
kernels/cuda_graph_runner.py:管理「桶 → 图」字典、惰性/预捕获、padding 与选桶。models.py:抽出 capture-safe 的_encode_and_predict(Transformer stack + 预测头);RepEncoder 留在图外(embedding 指针表是动态地址)。load_model()调用enable_cuda_graph() + precapture_all()一次性捕获全部桶,避免捕获开销落入计时循环。
全量测试集(2039 batch),bf16 + flash_triton + moe=triton:
| Eager | CUDA Graph | 提升 | |
|---|---|---|---|
| Inference time | 21.08s | 9.87s | 2.14×,省 11.2s |
| AUC | 0.757996 | 0.757994 | -2e-6(bf16 噪声) |
| PCOC | 1.111879 | 1.111884 | +5e-6 |
| score_all | 74.29 | 76.91 | +2.62 |
(该表为引入 CUDA Graph 时的历史测量,AUC 绝对值与当前 ckpt 略有差异。)
预测概率与 eager 路径在所有预测位上 bit-exact(padding 行不进入 pred_mask)。
build_env.sh 通过运行 infer.py 预热 Triton cache,将编译产物固化进 triton_cache/:
python infer.py # 跑一遍全量数据集,触发全部 kernel 的 JIT 编译- FlashAttn、Fused MoE 等 Triton kernel 在首次调用时 JIT 编译(~秒级)。
- CUDA Graph 的捕获过程也会触发全部 kernel 编译。
- 预热后提交评测时零编译开销。
python infer.py [--attn-mode flash_triton|sdpa] [--moe-backend triton|torch] \
[--cuda-graph|--no-cuda-graph] [--ckpt ckpt.pt]- 默认:
--attn-mode flash_triton --moe-backend triton --cuda-graph(最优配置)。 --attn-mode sdpa:PyTorch 标准 SDPA + 块对角因果掩码(大 batch 会 OOM)。--moe-backend torch:纯 PyTorch 实现的 SMoE(慢,但兼容性好)。--no-cuda-graph:关闭 CUDA Graph,回退 eager 执行(调试用)。
各优化阶段的累积收益(RTX 4090, bf16, 2039 batch, attn_mode=sdpa 基线):
| 优化阶段 | Latency | AUC | 加速比 |
|---|---|---|---|
| Eager baseline | 30.33s | 0.7577 | 1.00× |
| + Triton MoE & Fused Routing | 25.92s | 0.7577 | 1.17× |
| + Triton Batched Embedding | 23.36s | 0.7577 | 1.30× |
CUDA Graph 的单独收益(同一测试集,bf16 + flash_triton + moe=triton,见上节表格):
21.08s → 9.87s,2.14×。
本仓库当前默认配置实测(RTX 4090, 2039 batch):
inference time: 12.52s
AUC: 0.760070
PCOC: 1.113738
score_all: 76.45
延迟远低于 300s 红线,AUC/PCOC 均在 [0.65, 1.0] × [0.85, 1.15] 的安全区内。
评测机为 A800-80G,显存与算力均优于此处的 4090,实际延迟应不高于该值。
.
├── infer.py # 推理入口,数据加载 + 模型前向 + 打分
├── models.py # 模型定义:RepEncoder, TransformerEncoder, CTRModel
├── metrics.py # AUC/PCOC 计算与打分公式
├── local_env.py # 环境配置:Triton cache 路径设置
├── Utils/
│ ├── data_utils.py # 数据加载:load_sample_files, CTRUserDataset, collate_fn
│ └── profiler.py # 性能分析工具:record_scope, export_profiler_csv
├── kernels/
│ ├── batched_emb_triton.py # Triton 批量 EmbeddingBag 算子
│ ├── flash_attn.py # Triton FlashAttention 变长因果算子
│ ├── fused_moe_triton.py # Triton Fused MoE 算子(路由 + GEMM + 归约)
│ ├── fused_moe_cuda.py # CUDA 扩展:topk_softmax, moe_align_block_size
│ ├── cuda_graph_runner.py # 分桶 CUDA Graph 管理
│ └── csrc/
│ ├── topk_softmax.cu # CUDA kernel:softmax + top-k + renorm
│ └── moe_align.cu # CUDA kernel:专家对齐 + 路由准备
├── requirements.txt # 依赖:torch 2.6.0, triton 3.2.0, sklearn, tqdm
├── build_env.sh # 预热脚本:跑一遍推理,固化 Triton cache
└── ckpt.pt # 模型权重(10GB,不入库)
- FlashAttention-2: https://arxiv.org/abs/2307.08691
- Sparse Mixture-of-Experts: https://arxiv.org/abs/1701.06538
- Triton Language: https://github.com/openai/triton
- CUDA Graphs: https://docs.nvidia.com/cuda/cuda-c-programming-guide/index.html#cuda-graphs