🚌 第三课:融合算子替换 —— 源码 + 主图 + 逐算子拆解
⚡ 从“小车队”到“超级大巴” · 完整展现每个算子的 NPU 内部操作
源码层
PyTorch 调度
NPU 算子执行
融合优化
📂 第一步:源码 —— 未融合的 RMSNorm
transformers/models/qwen3/modeling_qwen3.py
Qwen3RMSNorm.forward 原始实现(8个小算子串行)
def forward(self, hidden_states):
input_dtype = hidden_states.dtype
hidden_states = hidden_states.to(torch.float32)
variance = hidden_states.pow(2)
variance = variance.mean(-1, keepdim=True)
variance = variance + self.variance_epsilon
rsqrt = torch.rsqrt(variance)
hidden_states = hidden_states * rsqrt
hidden_states = self.weight * hidden_states
return hidden_states.to(input_dtype)
🔢 共 8 个算子(含2次精度转换),每个都对应一次独立的 NPU 任务下发。
📊 主图:RMSNorm 融合前后完整调用链
flowchart TD
subgraph 业务代码层 ["🧑💻 业务模型代码层"]
A["Qwen3RMSNorm.forward
(transformers/models/qwen3/modeling_qwen3.py)"] --> B{选择执行路径}
end
subgraph PyTorch框架层 ["⚙️ PyTorch 框架层"]
B -->|未融合路径| C["调用基础算子
pow, mean, add, rsqrt, mul..."]
B -->|融合路径| D["调用融合算子入口
torch_npu.npu_rms_norm(...)"]
end
subgraph CANN算子层_未融合 ["📦 CANN 算子层 (未融合)"]
C --> C1["Pow 算子"]
C --> C2["Mean 算子"]
C --> C3["Add 算子"]
C --> C4["Rsqrt 算子"]
C --> C5["Mul 算子"]
C --> C6["Mul 算子"]
C1 & C2 & C3 & C4 & C5 & C6 --> C7["6次NPU任务下发
多次DDR读写"]
end
subgraph CANN算子层_融合 ["📦 CANN 算子层 (融合)"]
D --> D1["融合算子: RmsNorm"]
D1 --> D2["单次NPU任务下发
单次DDR读写"]
end
subgraph NPU硬件层 ["💻 NPU 硬件执行层"]
C7 --> E1["串行执行6个kernel
(含调度空泡)"]
D2 --> E2["单次执行全部计算
核内缓存流转"]
E1 --> F["返回结果"]
E2 --> F
end
style A fill:#e1f0fa,stroke:#1f5e8e
style C fill:#ffe6e6,stroke:#cc5555
style D fill:#e6ffe6,stroke:#2b8b5c
style D1 fill:#e6ffe6,stroke:#2b8b5c
style E2 fill:#e6ffe6,stroke:#2b8b5c,stroke-dasharray: 5 5
🔍 主图说明: 从源码出发,左侧红色分支为未融合路径(6个独立算子),右侧绿色分支为融合路径(1个融合算子)。最底层展示了 NPU 硬件执行的差异。
⚙️ 第三步:每个算子的 NPU 内部操作(独立图)
①.to(torch.float32) 精度转换 · 非计算密集型
flowchart LR
subgraph NPU["NPU AI Core"]
direction TB
A["读取 DDR
FP16 数据"] --> B["转换单元
FP16 → FP32"]
B --> C["写回 DDR
FP32 数据"]
end
D["CPU 调度"] -->|下发任务| NPU
文字示例: CPU 下发转换任务 → NPU 从 DDR 读取 FP16 张量 → 硬件转换单元逐元素转为 FP32 → 写回 DDR。耗时 ~0.02ms,产生 1 次 DDR 读写。
②.pow(2) 逐元素平方 · 计算密集型
flowchart LR
subgraph NPU["NPU AI Core"]
direction TB
A["读取 DDR
FP32 张量 x"] --> B["Vector 单元
x[i] = x[i] * x[i]"]
B --> C["写回 DDR
平方结果 x²"]
end
D["CPU 调度"] -->|下发任务| NPU
文字示例: CPU 下发 Pow 算子 → NPU Vector 单元逐元素平方(128 个元素/周期并行)→ 结果写回 DDR。耗时 ~0.15ms,产生 1 次读 + 1 次写。
③.mean(-1, keepdim=True) 规约求和 · 计算密集型
flowchart LR
subgraph NPU["NPU AI Core"]
direction TB
A["读取 DDR
x² 张量 [B,S,H]"] --> B["Vector 规约单元
沿 H 维求和 ÷ H"]
B --> C["写回 DDR
均值 [B,S,1]"]
end
D["CPU 调度"] -->|下发任务| NPU
文字示例: CPU 下发 Mean 算子 → NPU 读取平方结果 → 沿最后一维(H=1536)累加求和再除以 H → 形状从 [B,S,1536] 降为 [B,S,1] → 写回 DDR。耗时 ~0.20ms。
④+ self.variance_epsilon 逐元素加常数 · 计算轻量
flowchart LR
subgraph NPU["NPU AI Core"]
direction TB
A["读取 DDR
均值 [B,S,1]"] --> B["Vector 单元
逐元素 + 1e-6"]
B --> C["写回 DDR
variance+eps"]
end
D["CPU 调度"] -->|下发任务| NPU
文字示例: 读取均值结果 → 与常量 1e-6 逐元素相加 → 写回。耗时 ~0.01ms,但仍有 1 次读 + 1 次写。
⑤torch.rsqrt 1/√x · 计算密集型
flowchart LR
subgraph NPU["NPU AI Core"]
direction TB
A["读取 DDR
var+eps"] --> B["Vector 单元
1 / √(x)"]
B --> C["写回 DDR
rsqrt 结果"]
end
D["CPU 调度"] -->|下发任务| NPU
文字示例: 读取 var+eps → 调用硬件 rsqrt 指令(内置牛顿迭代法)→ 写回。耗时 ~0.12ms,需要多次迭代收敛。
⑥hidden_states * rsqrt 逐元素乘 · 计算密集型
flowchart LR
subgraph NPU["NPU AI Core"]
direction TB
A["读取 DDR
原始 x 和 rsqrt"] --> B["Vector 单元
x[i] * rsqrt[j](广播)"]
B --> C["写回 DDR
归一化结果"]
end
D["CPU 调度"] -->|下发任务| NPU
文字示例: 同时读取原始 x [B,S,H] 和 rsqrt [B,S,1] → 广播乘 → 写回。耗时 ~0.18ms,2 次读 + 1 次写。
⑦self.weight * hidden_states 逐元素乘权重 · 计算密集型
flowchart LR
subgraph NPU["NPU AI Core"]
direction TB
A["读取 DDR
归一化结果 和 weight"] --> B["Vector 单元
逐元素乘 weight"]
B --> C["写回 DDR
最终结果"]
end
D["CPU 调度"] -->|下发任务| NPU
文字示例: 读取归一化结果和权重 → 逐元素乘 → 写回。耗时 ~0.18ms,2 次读 + 1 次写。
⑧.to(input_dtype) 精度转换 · 非计算密集型
flowchart LR
subgraph NPU["NPU AI Core"]
direction TB
A["读取 DDR
FP32 结果"] --> B["转换单元
FP32 → FP16"]
B --> C["写回 DDR
FP16 结果"]
end
D["CPU 调度"] -->|下发任务| NPU
文字示例: 读取 FP32 结果 → 转换为 FP16 → 写回。耗时 ~0.02ms,产生 1 次读 + 1 次写。
⏳ 第四步:6次NPU任务下发 + 调度空泡(Timeline)
⚠️ 关键问题: 每个算子都需要 CPU 单独下发任务,NPU 在执行完一个算子后,必须 等待 CPU 下发下一个任务,这期间 NPU 处于空闲状态 —— 这就是 调度空泡(Scheduling Bubble)。
📊 未融合模式 —— NPU 时间线(含调度空泡)
时间轴 → 0ms 0.5ms 1.0ms 1.5ms 2.0ms
─────────────────────────────────────────────────────────────
CPU调度: [下发①] [空闲] [下发②] [下发③] [下发④] [下发⑤] [下发⑥]
█████████ █████████ █████████ █████████ █████████
NPU执行: [ ① ] [ ② ] [ ③ ] [ ④ ] [ ⑤ ] [ ⑥ ]
████████ ████████ ████████ ████████ ████████
↑ 执行 ↑ ↑ 执行 ↑ ↑ 执行 ↑ ↑ 执行 ↑ ↑ 执行 ↑ ↑ 执行 ↑
██████████ ██████████ ██████████ ██████████ ██████████
← 调度空泡 → ← 调度空泡 → ← 调度空泡 → ← 调度空泡 → ← 调度空泡 →
🔴 红色 = NPU 空闲(调度空泡) | 🟢 绿色 = NPU 执行 | 🟠 CPU调度
📌 总耗时: 6 个算子各自执行耗时(0.15+0.20+0.01+0.12+0.18+0.18 = 0.84ms)+ 5 次调度空泡(每次 ~0.1ms = 0.5ms) = ~1.34ms
🚀 第五步:融合算子 —— torch_npu.npu_rms_norm
替换后的 Qwen3RMSNorm.forward
def fused_forward(self, hidden_states):
return torch_npu.npu_rms_norm(
hidden_states,
self.weight,
self.variance_epsilon
)[0]
📊 融合算子 —— NPU 内部完整操作图
flowchart TD
subgraph CPU调度["CPU 调度层"]
A["单次下发任务
torch_npu.npu_rms_norm"]
end
subgraph NPU内部["NPU AI Core 内部流水线"]
direction TB
B["Vector 单元
逐元素平方"] --> C["Vector 单元
沿 H 维求均值"]
C --> D["标量加法
+ epsilon"]
D --> E["Vector 单元
rsqrt 指令"]
E --> F["Vector 单元
广播乘(归一化)"]
F --> G["Vector 单元
逐元素乘权重"]
G --> H["写回 DDR
FP16 最终结果"]
end
A -->|一次 DMA 读| B
H -->|一次 DMA 写| I["返回结果"]
style A fill:#d9e6f5,stroke:#2b7ba8
style B fill:#e6ffe6,stroke:#1f8b6b
style C fill:#e6ffe6,stroke:#1f8b6b
style D fill:#e6ffe6,stroke:#1f8b6b
style E fill:#e6ffe6,stroke:#1f8b6b
style F fill:#e6ffe6,stroke:#1f8b6b
style G fill:#e6ffe6,stroke:#1f8b6b
style H fill:#e6ffe6,stroke:#1f8b6b
🔑 关键差异: 中间结果 全部在 NPU 核内缓存(L1/L2)中流转,不写回 DDR,直到最终结果才写回一次。
📊 融合模式 —— NPU 时间线(无调度空泡)
时间轴 → 0ms 0.5ms 1.0ms
─────────────────────────────────────────
CPU调度: [下发融合任务] [空闲] (等待完成)
████████████████
NPU执行: [ 融合算子连续执行 平方→均值→加eps→rsqrt→归一化→乘权重 ]
████████████████████████████████
← 无调度空泡,NPU 满负荷运转 →
📌 总耗时: 融合算子内部计算耗时(平方+均值+加eps+rsqrt+归一化+乘权重 ≈ 0.84ms)+ 0 次调度空泡 = ~0.84ms
🚀 加速比: 1.34ms / 0.84ms ≈ 1.6x(仅替换 RMSNorm 一个单元)
📊 融合前后完整对比
| 对比维度 | 未融合(6个小算子) | 融合(npu_rms_norm) |
| CPU 下发次数 | 6 次 | 1 次 |
| DDR 读写次数 | ~11 次(每算子读+写) | 2 次(1次读输入 + 1次写结果) |
| 调度空泡 | 5 次(每次 ~0.1ms) | 0 次 |
| 中间结果存储 | DDR(5次中间写回) | 核内缓存(L1/L2) |
| 总耗时(6算子) | ~1.34ms | ~0.84ms |
| 全模型影响 | 57 个 RMSNorm × 6 次下发 = 342 次 | 57 次下发 |
💡 核心结论: 融合算子通过 减少 CPU 下发次数(6→1)、消除调度空泡、减少 DDR 读写,在 NPU 上实现了 1.6x 的纯算子层加速。与图编译叠加后效果更显著。
📖 基于 CANN 学习中心 · 第三课 · 源码+主图+逐算子拆解版