Skip to content
 
 

Repository files navigation

CTR 模型前向推理流程

本仓库实现了一个基于 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

前向传播详细流程

输入:Batch 字典

键 类型 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),每个用户内部保持时序因果。


阶段 1:RepEncoder(特征聚合)

输入:batch 字典
输出:seq_input → [total_tokens, 512]

操作流程:

  1. Batched Embedding Bag(Triton 融合算子):

    • 对 28 个槽位并行调用 triton_batched_embedding_bag,一次 kernel launch 完成全部查表。
    • 每个槽位内的多个特征 ID 做 sum pooling(mode="sum")。
    • 输出 [batch_size, 28, 512] → reshape 为 [batch_size, 14336]。
  2. 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× 加速)。

阶段 2:TransformerEncoder(序列建模)

输入: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

2.1 FlashAttention Varlen (Triton 自研)

为什么需要:

  • 打包序列是 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 行只看 key 0..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× 加速)

2.2 Sparse Mixture-of-Experts (Triton 融合)

结构:

  • 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):

  1. 路由准备(GPU 异步):

    sorted_route_ids, expert_ids, num_tokens_post_pad = prepare_routing_triton(...)
    # - 用 scatter_add 统计每个专家分到的 token 数
    # - CUDA kernel 一次完成:专家 ID 填充 + token 重排 + route_positions 记录
  2. 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
  3. W2 矩阵乘 + scatter 输出:

    route_out = grouped_matmul_scatter_routes_triton(hidden, w2, ...)
    # 直接写入 route-major layout:row `token_id * top_k + k_idx`
  4. 加权归约:

    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 开销。

阶段 3:Prediction Head

输入: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 输出。


分桶 CUDA Graph 优化

为什么需要

瓶颈:每个 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 之间变化。

解决:

  1. 分桶 padding:以 512 为步长设 10 个桶(2560, 3072, ..., 7168),覆盖 98.6% batch。

    • 每个桶惰性捕获一张独立的图,作用在固定 [S_PAD, 512] 的静态 buffer 上。
    • 超过顶桶(1.4% 长尾)回退 eager,避免极少数大 batch 撑爆显存。
  2. 固定 attention grid:

    • 捕获时用固定 max_seqlen = S_PAD 决定 FlashAttn 的 grid 维度。
    • Triton kernel 的逐块 early-exit(q_block_start >= seqlen 直接 return)让"偏大的 grid"对真实较短的段成为空操作,结果完全不变。
  3. 去除 host-sync:

    • 原 max_seqlen = int(lengths.max().item()) 的 .item() 会强制同步,无法进图。
    • 分桶后每桶 max_seqlen 是常量,直接传入。
  4. 变长 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)。


部署提示

预热 Triton Cache

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,不入库)

参考资料

About

No description, website, or topics provided.

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages