副标题:同一个算子,训练和推理在硬件眼里是两种完全不同的动物。Conv2D、GEMM、Normalization、Attention——每个算子都在训练时和推理时面临不同的 shape、精度、访存和融合需求。本文从算子开发工程师的视角,系统性地拆解这些差异。
引子:一个 Conv2D 算子的两面
先从一个最简单的例子说起。
Conv2D——深度学习最基础的算子。不管是什么 AI 芯片,都要实现它。但同一个 Conv2D,在训练和推理时的实现路径可以完全不同:
| 训练 | 推理 | |
|---|---|---|
| Batch size | 64-512(大) | 1-8(小) |
| 数据类型 | FP32 / BF16 | FP16 / INT8 / FP8 |
| 需要什么 | Forward + Backward ×2 | Forward 仅 |
| 核心算法偏好 | Winograd(乘少换快) | Implicit GEMM(无缓存) |
| 融合模式 | Conv + BNUpdate + ReLU | Conv + Bias + ReLU |
前向卷积的数学是一样的——不同在于:算子的 shape、精度要求、融合策略和是否需要反向传播,把同一个算子的硬件实现拉到了两条不同的优化路径上。
这不是 Conv2D 独有的。GEMM、Normalization、Attention 每一个算子都在训练和推理之间展现出了同样深刻的差异。理解这些差异,对自研芯片的算子开发有决定性的意义。
一、Conv2D:四家厂商的路线分歧
1.1 四种核心算法
Conv2D 的核心挑战是滑动窗口访存不连续。所有芯片厂商的实现都在以下四种算法中选择:
| 算法 | 核心思想 | 乘法复杂度 | 显存额外开销 | 适用场景 |
|---|---|---|---|---|
| Direct | 7 层循环硬算 | O(N·C_out·H_out·W_out·C_in·K²) | 0 | 几乎不用,效率 < 1% |
| Im2Col + GEMM | 展开成矩阵做 GEMM | 同上 | K² 倍输入大小 | 有矩阵乘单元的芯片 |
| Implicit GEMM | GEMM kernel 内在线算地址 | 同上 | 0(不显式展开) | 当前主流 |
| Winograd | 数学变换减少乘法 | ~44%(K=3 时) | 变换矩阵中间量 | 3×3、大 batch |
| FFT | FFT→点乘→IFFT | O(N log N) | 频域缓存 | 大 kernel(≥7×7) |
1.2 四家厂商的路线选择
| 厂商 | 核心路线 | 为什么走这条路 |
|---|---|---|
| NVIDIA cuDNN | Implicit GEMM + Tensor Core | Tensor Core 只擅长 GEMM,Implicit GEMM 把卷积变成 GEMM 同时省掉 im2col 缓存。Find API 在多种算法间自动选最优。 |
| AMD MIOpen | Implicit GEMM +汇编 Winograd | Matrix Core 同样适合 GEMM。Winograd 用纯汇编手写(没有 C++ 版本),利用 CDNA 的 64-lane wavefront 做高效 shuffle。 |
| Intel oneDNN | JIT 编译生成 | 不手写多个 kernel 变体,而是运行时按参数生成刚好匹配的 kernel。CPU 上用 AVX-512/AMX 向量化,GPU 上用 DPAS。 |
| 华为昇腾 CANN | 显式 Im2Col + Cube GEMM | 达芬奇架构的 Cube Unit不能做 Implicit GEMM(无法在矩阵乘内部做条件地址计算),只能用 Vector Unit 做显式 im2col 展开,再喂给 Cube。 |
1.3 训练和推理在 Conv2D 上的差异
训练选 Winograd,推理选 Implicit GEMM:
训练(N=128, 3×3 Conv): Winograd 把乘法从 9 降到 4 N=128 分摊了变换矩阵的额外计算 总 FLOPs 减少 ~50% → Winograd 胜出 推理(N=1, 3×3 Conv): Winograd 的变换步骤(A^T·X·A, G·W·G^T)多了 3 次小矩阵乘 在 N=1 时,变换开销可能超过省下的乘法 → Implicit GEMM 胜出(无缓存、无变换、1 个 kernel 搞定)训练必须 FP32 累加,推理可以用 INT8:
训练时正向的微小误差会被反向传播放大。因此即使输入是 FP16,训练卷积的累加器(accumulator)也必须是 FP32。推理没有反向传播,INT8 量化卷积的精度损失可以接受。
训练融合 BN update,推理吸收 BN:
训练融合: Conv2D + BN(update) + ReLU → 融合 kernel 因为 BN 前向要更新 running_mean/var,不能直接吸收到 Conv 权重 推理: BN(γ·x+β) 是线性变换,直接吸收进 Conv2D 的 W' 和 b' W' = γ·W / sqrt(σ²+ε) b' = γ·b / sqrt(σ²+ε) + β - γ·μ / sqrt(σ²+ε) → Conv2D 融合 kernel 不需要 BN 参与,更简单训练有反向算子,推理没有:
训练需要额外实现Conv2DBackpropInput(计算输入梯度)和Conv2DBackpropFilter(计算权重梯度)。这两个反向算子的计算模式不同:BackpropInput等价于转置卷积,BackpropFilter需要另一种数据流。
Conv2D 的差异在所有算子中是最明显的——因为它的访存模式特殊,算法选择多。接下来的 GEMM、Normalization、Attention 的差异更隐蔽,但影响同样深远。
二、GEMM:训练要吞吐,推理要延迟
GEMM(C = A @ B)是 AI 推理训练中占比最大的算子(LLM 中占 ~65% 算力)。但训练和推理场景下的 GEMM 几乎是两个世界。
2.1 Shape 差异:天壤之别
训练 GEMM 的 shape:
Forward: Y = X @ W^T → M = B × L, N = d_out, K = d_in Backward: dX = dY @ W → M = B × L, N = d_in, K = d_out dW = dY^T @ X → M = d_out, N = d_in, K = B × L 典型值: B=64, L=2048, d_in=4096, d_out=4096 → M = 131072, N = 4096, K = 4096 → 巨大的 GEMM,算力打满推理 Decode GEMM 的 shape:
Decode: Y = X @ W^T → M = B, N = d_out, K = d_in 典型值: B=1, d_in=7168, d_out=4096 → M = 1, N = 4096, K = 7168 → GEMV(矩阵×向量),算力利用率极低这个差异意味着什么:
训练 GEMM(M=131072): Arithmetic Intensity ≈ M·N·K / (M·K + N·K + M·N) · dtype_size ≈ 131072·4096·4096 / (131072·4096 + 4096·4096 + 131072·4096) · 2 ≈ 4.4 TFLOPs / 1.1 GB = ~4000 FLOP/Byte → 深度计算密集,可以轻松打满 Tensor Core 推理 Decode GEMM(M=1): Arithmetic Intensity ≈ 1·4096·7168 / (1·7168 + 4096·7168 + 1·4096) · 2 ≈ 29M / 29.4MB = ~1 FLOP/Byte → 极度访存密集,带宽是唯一瓶颈2.2 精度差异
| 训练 | 推理 | |
|---|---|---|
| 输入精度 | BF16 / FP16 | FP16 / INT8 / FP8 |
| 累加精度 | 必须 FP32 | FP16 可以接受 |
| 权重精度 | FP32 / BF16 | FP16 / INT4 / FP8 |
为什么训练必须 FP32 累加?
FP16 的最大表示范围约 65504,GEMM 累加时中间结果很容易超过这个范围。一个 4096×4096 的 GEMM,K=4096 次累加,每次累加的是 FP16 乘积(值域 ~ ±65504²),平均可能达到 10⁹ 的量级——远超 FP16 的表示范围。如果不使用 FP32 累加器,精度直接溢出到 NaN,梯度全丢。
推理中不存在反向传播,累加误差最多影响一个输出 token,不会被放大。
2.3 算法差异
训练 GEMM(M >> 1): Tile GEMM: 把大矩阵切成 tile(如 128×128),每个 tile 对应一个 thread block 在 Tensor Core 上运行,共享内存做 double buffering 目标是打满 Tensor Core 利用率(H100 可达 ~80% 理论峰值) 推理 Decode GEMM(M = 1): GEMV 优化: M=1 时不能用标准的 tile GEMM(tile 高度为 128,但 M=1,利用率 < 1%) 有两种方案: A) Split-K: 把 K 维度切分成多份并行计算,最后 reduce B) Warp-level GEMV: 一个 warp 处理所有 N 维度,逐 K 步进 目标是最大化带宽利用率(达到 ~80% HBM 带宽即可)对自研芯片的影响:
如果你的芯片只为训练优化(大 tile GEMM),推理 decode 的 GEMV 场景下性能会暴跌。必须同时有"大 GEMM kernel"和"小 GEMV kernel",或者统一的 kernel 能自适应 M 维度。
2.4 一个特殊的精度问题:确定性(Determinism)
训练场景有一个推理没有的要求——多批次之间必须 bitwise 一致:
训练: batch 64 和 batch 128 跑同一个输入 → 前 64 个样本的梯度必须完全一样 否则调 bug 时无法复现,分布式训练时更新不一致 → 需要确定性 GEMM(sgemm-bi 的做法): 固定累加顺序,不用原子操作,不同 batch 之间结果 bit-identical 推理: 不需要确定性,一次 Forward 就够了 不同的推理步骤可以有不同的精度三、Normalization:训练要统计,推理要固定
3.1 BatchNorm:差异最大
BatchNorm 是训练和推理行为差异最大的算子:
训练时: 输入 x (N, C, H, W) ↓ 计算当前 batch 的 μ = mean(x), σ² = var(x) ↓ y = γ · (x - μ) / sqrt(σ² + ε) + β ↓ 更新 running_mean = momentum × running_mean + (1-momentum) × μ 更新 running_var = momentum × running_var + (1-momentum) × σ² → μ 和 σ² 每次都不一样(取决于当前 batch 的数据) 推理时: 输入 x (N, C, H, W) ↓ y = γ · (x - fixed_mean) / sqrt(fixed_var + ε) + β → 使用训练中积累的固定 running_mean/running_var → μ 和 σ² 是固定的,没有数据依赖性对算子实现的直接影响:
训练 BN: 需要 2 次规约(求 mean 和 var)+ 2 次更新(running stats) 需要存中间输入 x 用于反向传播 → 算子实现更重,有额外的统计量计算 推理 BN: 一次逐元素乘加(线性变换)就够了 → 实践中可以直接吸收进前一个 Conv2D 的权重里 → 根本没有独立的 BN 算子BN 融合进 Conv2D(推理-only 优化):
推理时: Conv2D(x, W, b) → BN(x') = γ·x' + β 合并: Conv2D(x, W', b') W' = γ·W / sqrt(σ²+ε) b' = γ·(b - μ) / sqrt(σ²+ε) + β 省: 一次 HBM 读写(不需要存 x' 中间结果)从这个意义上说,推理场景的"BN"根本不是一个算子——它被优化没了。
3.2 LayerNorm / RMSNorm:训练存中间量,推理即算即用
相比 BN,LayerNorm 和 RMSNorm 的训练/推理差异小得多——因为它们不存在"数据依赖的统计量"。但差异仍然存在:
训练 LayerNorm: y = γ·(x - μ) / sqrt(σ² + ε) + β 需要存 x 和 σ² 用于反向传播 → 显存占用: 每层多存 1 个 x 的副本 (d_model,) 推理 LayerNorm: 不需要存任何中间结果 算完 y 直接用,x 的内存可以立即释放LLM 中使用的是 RMSNorm(比 LayerNorm 少算 μ 和 β)。在推理场景下 RMSNorm几乎总是和前面的 GEMM/注意力输出融合——作为 GEMM epilogue 的一部分,直接在寄存器中完成归一化,不进 HBM。
3.3 厂商实现差异
| 厂商 | 训练 Normalization | 推理 Normalization |
|---|---|---|
| NVIDIA | cuDNN BN 有训练/推理两个 mode(CUDNN_BATCHNORM_TRAINING/CUDNN_BATCHNORM_INFERENCE) | 推理 mode 不更新 running stats |
| Intel oneDNN | dnnl::batch_normalization_forward有training和inferenceprop_kind | 推理跳过 stats 更新 |
| 华为 CANN | 训练特有融合:BNTrainingUpdate + Conv2D + BNTrainingReduce | 推理直接吸收 BN 到 Conv |
四、Attention:计算模式完全不同的两个算子
如果 Conv2D 和 GEMM 还是"同一个算子、不同参数"的话,Attention 在训练和推理之间的差异已经到了可以视为两个不同算子的程度。
4.1 训练:并行处理所有 token
训练 Attention(以 MHA 为例): Q = x @ W_q # (B×L, h×d_k) ← 所有 token 并行 K = x @ W_k # (B×L, h×d_k) ← 所有 token 并行 V = x @ W_v # (B×L, h×d_k) ← 所有 token 并行 score = Q @ K^T # (B×L, B×L) ← 巨大的注意力矩阵! O(n²) 的核心开销 attn = softmax(score / √d) out = attn @ V # (B×L, h×d_v) 反向传播: 需要存 Q、K、V、score、attn、out、x 全部中间量 显存占比: ~70% 的总训练显存花在存 attention 中间激活上关键特征:
- 计算密集:Q@K^T 是大 GEMM(M=L, N=L, K=d_k),打满 Tensor Core
- 显存爆炸:L 增大时中间激活 O(L²) 增长
- 反向需要所有中间量,催生了 FlashAttention 和 Gradient Checkpointing
4.2 推理 Decode:逐 token 自回归
推理 Decode(生成第 t+1 个 token): q_t+1 = x_t+1 @ W_q # (B, h×d_k) ← 只算当前 token k_t+1 = x_t+1 @ W_k # (B, h×d_k) ← 只算当前 token v_t+1 = x_t+1 @ W_v # (B, h×d_k) ← 只算当前 token # KV Cache: 把之前所有 token 的 K/V 存起来 K_cache = concat(K_cache, k_t+1) # (B, L+1, h×d_k) V_cache = concat(V_cache, v_t+1) # (B, L+1, h×d_k) # Attention with cache score = q_t+1 @ K_cache^T # (B, 1, h) @ (B, L, h) → (B, 1, L) ← O(n) 不是 O(n²)! attn_out = score @ V_cache # (B, 1, L) @ (B, L, h) → (B, 1, h)关键特征:
- 访存密集:瓶颈在读取 KV Cache(B×L×d_k×2 bytes/步)
- O(n) 不是 O(n²):因为只算了 q_t+1 和所有 K 的 score,没有算 Q 全矩阵
- KV Cache 线性增长:长度 L 增加,每步多存 2×d_k×dtype bytes
4.3 一张表看清差异
| 维度 | 训练 | 推理 Decode |
|---|---|---|
| 计算模式 | 所有 token 并行 | 逐 token 自回归 |
| 核心计算 | Q @ K^T (L×L×d_k) | q @ K_cache^T (L×d_k) |
| 复杂度 | O(L²) | O(L) |
| 瓶颈 | 计算(Tensor Core 利用率) | 访存(KV Cache 带宽) |
| KV Cache | 不需要 | 必需——每步增长 |
| 中间激活 | 存 Q/K/V/score/attn 全部 | 不存任何中间量 |
| FlashAttention 角色 | 消除 O(L²) 的 HBM 读写 | 优化 KV Cache 加载 |
4.4 FlashAttention 在训练和推理中的不同角色
这是“同一个优化技术、面对完全不同的瓶颈”的典型案例:
训练(FlashAttention tiling): 问题: Q@K^T 的中间结果 (L×L) 太大,写 HBM→读 HBM 浪费带宽 解法: 分块 tiling——在共享内存里算完 attention,只写回最终结果 收益: I/O 复杂度从 O(L²+L) 降到 O(L²/√M),显存省了 10x 效果: 训练速度 2-3x,大 batch 下也能跑长序列 推理(FlashAttention 解码): 问题: KV Cache 太大,每步要加载全部 K/V 到片上 解法: 多个 thread block 并行加载 KV Cache 的不同分块 收益: 饱和 HBM 带宽,不浪费 效果: Decode 加速最高 28x 同一个 FlashAttention kernel, 训练时解决的瓶颈是"中间矩阵写 HBM 太慢", 推理时解决的瓶颈是"KV Cache 从 HBM 加载太慢"。4.5 对自研芯片的意义
Attention 在训练和推理之间的差异是所有算子中最大的。如果你的芯片只优化了训练的"大 GEMM + tiled softmax",推理时可能会遇到:
- GEMV 瓶颈:q @ K^T 在 decode 时是 1×d_k @ d_k×L,一个大 GEMV。如果你的芯片没优化 GEMV,Tensor Core 利用率接近 0%。
- KV Cache 带宽瓶颈:每步要读 L×d_k 的数据。如果 HBM 不够宽,decode 延迟直接卡住。
- PagedAttention 支持:推理需要用 PagedAttention 来管理 KV Cache 显存。这不是"提速",这是"能不能跑"(没有分页,64K 上下文的 KV Cache 浪费 60%+ 显存)。
五、激活函数:差异最小
激活函数(ReLU、SiLU、GELU、SiTU)是训练/推理差异最小的算子。前向计算完全一致。唯一的差异在反向:
训练: 需要存输入 x → 用于反向计算梯度 ReLU: dL/dx = dL/dy · (x > 0) ← 需要 x SiLU: dL/dx = dL/dy · [sigmoid(x) + x·sigmoid(x)·(1-sigmoid(x))] ← 需要 x 如果不存 x,就要重新算前向,浪费算力 推理: 只算前向: y = f(x) 不需要存任何东西 输出可以直接覆盖输入,省显存对实现的启示:激活函数本身不需要区分训练/推理版本。差异在内存管理策略——训练需要分配额外的 buffer 存输入,推理不需要。
六、系统级差异:贯穿所有算子的主线
以上逐算子的分析背后,有几条贯穿所有算子的主线:
6.1 反向传播的存在决定了显存策略
训练: Forward 需要保存中间激活 → 反向才能算梯度 通常每层要存 ~4x 的中间量(输入、输出、注意力矩阵等) 显存大头 = 中间激活(不是权重!) 推理: Forward 算完即弃 显存大头 = 权重 + KV Cache这对芯片设计的含义——训练芯片需要更大的 HBM 容量(存中间激活)和计算/显存平衡的设计;推理芯片的 HBM 可以主要分配给权重和 KV Cache。
6.2 Batch size 决定了计算访存特征
| 算子维度 | 训练 | 推理 |
|---|---|---|
| Batch size | 64-512 | 1-8 |
| M 维度 | 大(B×L) | 小(B) |
| 瓶颈 | 计算(FLOPs 打满) | 访存(HBM 带宽) |
| Arithmetic Intensity | 1000-10000 FLOP/Byte | 0.5-5 FLOP/Byte |
训练芯片需要强大的算力(TFLOPS),推理芯片需要强大的带宽(TB/s)。这两者是不同的设计目标——H100 有 3.35 TB/s 和 1979 TFLOPS,H200 把带宽提到 4.8 TB/s 但算力没变,B200 提到 8 TB/s。NVIDIA 一直在加带宽,因为推理场景带宽比算力更值钱。
6.3 融合策略不同
| 算子 | 训练融合 | 推理融合 |
|---|---|---|
| Conv2D | Conv + BN(update) + ReLU | Conv + Bias + ReLU |
| GEMM | 一般不融合 | GEMM + Bias + Activation |
| Norm | 单独算(需要反向) | 吸收进前序 GEMM/Conv |
| Attention | FlashAttention(tile 融合) | PagedAttention + KV Cache |
推理的融合可以更大胆——因为没有反向传播的限制,可以把更多连续操作合并进一个 kernel,节省 HBM 读写。
6.4 自动调优策略不同
| 训练 | 推理 | |
|---|---|---|
| Find API 开销 | 可以接受(几十秒 vs 几小时训练) | 不能接受(首次加载必须快) |
| 算法选择 | 大 batch 偏好 Winograd/Tiled GEMM | 小 batch 偏好 Implicit GEMM/GEMV |
| 缓存策略 | 训练过程中 shape 固定,缓存一直有效 | shape 多变,可能需要规则引擎 |
| 确定性 | 需要(bitwise 可复现) | 不需要 |
七、实操建议:对算子开发团队的参考
7.1 如果只能先做一个,做推理
当前国内 AI 芯片的主要落地场景是推理部署。理由:
- Kernel 数量少:推理只需要 Forward,不需要 Backward。一个卷积算子是 Conv2D 不是 Conv2D + Conv2DBpropInput + Conv2DBpropFilter,工作量少 2/3。
- 精度要求低:INT8/FP16 推理是标准做法,不需要 FP32 累加器的大面积设计。
- 更容易做融合:不需要考虑反向传播的中间量保留,可以大胆做算子融合。
- 市场更成熟:云端推理、端侧推理、边缘推理都有明确需求。
7.2 推理优先,但 kernel 设计要为训练预留
虽然优先做推理,但 kernel 的接口设计应该为训练预留:
// 建议的卷积 API(训练和推理共用 forward):typedefstruct{intbatch_size;intinput_c,input_h,input_w;intoutput_c,kernel_h,kernel_w;intstride_h,stride_w;intpad_h,pad_w;DataType input_dtype;DataType weight_dtype;DataType accum_dtype;// FP32 for training, FP16 for inferencebool need_activation;// true: fused ReLUbool need_save_intermediate;// true: save input for backward (training)}Conv2DDesc;差别就在accum_dtype(FP32 还是 FP16)和need_save_intermediate(存不存中间量)两个参数上——kernel 的数学逻辑是一样的。
7.3 统一算子 + 两套调度策略
不要为训练和推理写两套 kernel。改用一套 kernel + 两套调用参数:
统一 GEMM kernel: gemm(M, N, K, A, B, C, accum_dtype, split_k) 训练调用: gemm(M=131072, N=4096, K=4096, accum=fp32, split_k=1) 推理调用: gemm(M=1, N=4096, K=7168, accum=fp16, split_k=8) ↑ 同一套代码,用 split_k 处理小 M 场景- M 大时自动走 tile GEMM 路径
- M 小时自动走 Split-K / GEMV 路径
accum_dtype控制累加器宽度
7.4 精度管理的通用策略
| 场景 | 输入/权重精度 | 累加器精度 | 说明 |
|---|---|---|---|
| 训练 Forward | FP16/BF16 | FP32 | 防止累加溢出 |
| 训练 Backward | FP16/BF16 | FP32 | 梯度精度更敏感 |
| 推理(FP16) | FP16 | FP16 或 FP32 | 可选,FP32更安全 |
| 推理(INT8) | INT8 | INT32 | 8 位乘 8 位累加用 32 位 |
| 推理(FP8) | FP8 | FP16/FP32 | FP8 动态范围窄,必须高精度累加 |
一条原则:累加器 = 2 × max(dtype bits)。8-bit 输入用 16-bit 累加不安全,必须 32-bit。16-bit 输入用 32-bit 累加。32-bit 输入用 32-bit 累加。
八、总结
核心发现
训练和推理的算子差异不在数学定义上,而在 shape、精度、反向传播和融合策略上。同样的 Conv2D、GEMM 和 Attention,在训练和推理场景下是两套完全不同的优化路径。
Attention 的差异最大——训练是 O(L²) 的大 GEMM + tiled softmax,推理是 O(L) 的 GEMV + KV Cache 加载。可以说是同一个名字的两个不同算子。
Batch size 是差异的根本来源——训练 M=B×L >> 1,所有算子是计算密集的;推理 M=B(通常 1-8),所有算子是访存密集的。
训练芯片和推理芯片的设计目标不同——训练需要峰值算力(TFLOPS),推理需要高带宽(TB/s)。NVIDIA 的 H100(3.35 TB/s)和 H200(4.8 TB/s)说明了这个趋势。
不要为训练和推理写两套 kernel——用统一 kernel + 两套调用参数(accum_dtype、split_k 配置)。数学逻辑一样,调度策略不同。
跨算子的差异矩阵
| 算子 | 训练 M 维度 | 推理 M 维度 | 训练精度 | 推理精度 | 训练额外算子 | 推理特有优化 |
|---|---|---|---|---|---|---|
| Conv2D | B×L(大) | B(小) | FP32 累加 | INT8 可 | Backprop ×2 | BN 吸收 |
| GEMM | B×L(大) | B(小) | FP32 累加 | INT8 可 | Backprop ×2 | Split-K GEMV |
| LayerNorm | B×L(大) | B(小) | FP32 | FP16 | 存 x 用于反向 | 融合进 GEMM epilogue |
| BatchNorm | B(大) | B(小) | FP32 | FP32 | running stats 更新 | 吸收进 Conv |
| Attention | L(长) | L(长) | FP32 累加 | FP16 | 存全中间量 | KV Cache + PagedAttn |
| ReLU/SiLU | B×L | B | FP32 | FP16/INT8 | 存 x 用于反向 | 融合进 GEMM epilogue |
附录:进一步阅读
- NVIDIA cuDNN 开发指南: docs.nvidia.com
- AMD MIOpen: github.com/ROCm/MIOpen
- Intel oneDNN: oneapi-src.github.io/oneDNN
- 华为 CANN 算子开发: hiascend.com
- FlashAttention: Dao et al. (2022):FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness
- sgemm-bi: Deterministic, batch-invariant GEMM for training