作者​:昇腾实战派
知识地图​:https://blog.csdn.net/Lumos_Lovegood/article/details/161601003

背景概述

GELU(Gaussian Error Linear Unit)激活函数因其平滑的梯度特性,在Transformer、BERT等现代深度学习模型中得到了广泛应用。在昇腾AI处理器的算子开发中,选择合适的编程模式对于平衡开发效率和运行性能至关重要。本文基于实际开发经验,详细介绍了采用SIMT(Single Instruction Multiple Threads)编程模式实现GELU算子的完整过程,包括设计规格、编程模型、Kernel实现、Host侧调度、精度验证及性能优化等关键环节,为开发者提供一套可参考的实践方案。

AscendC算子开发–SIMT GELU

1. 算子概述

1.1 功能描述

GELU(Gaussian Error Linear Unit)是一种常用的神经网络激活函数,相比 ReLU 具有更平滑的梯度特性,广泛应用于 Transformer、BERT 等现代网络架构中。

本算子采用 SIMT(Single Instruction Multiple Threads) 编程模式实现,每个线程独立处理一个元素,天然支持任意 shape、任意 axis 的计算需求。

1.2 计算公式

GELU 近似计算公式(tanh 近似展开):

G E L U ( x ) ≈ x 1 + e − 1.595769 ⋅ ( x + 0.044715 ⋅ x 3 ) GELU(x) \approx \frac{x}{1 + e^{-1.595769 \cdot (x + 0.044715 \cdot x^3)}} GELU(x)1+e1.595769(x+0.044715x3)x

其中:

  • − 1.595769 = − 2 ⋅ 2 π -1.595769 = -2 \cdot \sqrt{\frac{2}{\pi}} 1.595769=2π2 ,线性项系数
  • 0.044715 0.044715 0.044715,立方项原始系数

1.3 编程模式选择

维度 SIMD SIMT
调度单元 向量(一次处理多个元素) 线程(每个线程处理一个元素)
控制流 所有通道执行相同指令 每个线程有独立控制流
访存模式 要求连续对齐 支持随机访存
适合场景 规则计算、大批量连续数据 控制流复杂、访存不规则
GELU 适用性 ✅ 适合(纯逐元素计算,访存连续) ✅ 适合(编程模型简单直观)

SIMT 模式适合 GELU 的原因:

  1. GELU 是纯逐元素计算,每个线程独立处理一个元素,无数据依赖
  2. SIMT 编程模型与 CUDA 风格一致,开发者学习成本低
  3. 支持任意 shape,无需手动编写 tiling 逻辑

2. 设计规格

2.1 输入/输出定义

参数 Shape Data Type Format 说明
x(输入) 任意 shape float / half ND 输入张量
y(输出) 与 x 相同 float / half ND 输出张量

2.2 规格限制

限制项 约束值 说明
总元素数 ≤ 2³² - 1 uint32_t 索引上限
线程块大小 ≤ 2048 Ascend 950 AIV 硬件限制
Grid 线程块总数 ≤ 65535 Ascend 950 硬件限制
UB 总大小 256KB 每个 AIV 的片上内存

2.3 数据类型支持

输入类型 输出类型 说明
float float 标准模式
half float half_to_float 模式
half half 标准 half 模式
float half 降精度模式(可选)

3. 编程模型设计

3.1 线程组织

采用一维线程组织方式:

全局线程索引: thread_idx = blockIdx.x * blockDim.x + threadIdx.x
每个线程处理一个元素: y[thread_idx] = gelu(x[thread_idx])

线程调度策略:

  • 优先按 AIV 核数分配 block_num,充分利用硬件并行能力
  • 每个 block 内线程数取 32 的整数倍(warp 对齐),避免最后一个 warp 存在空闲通道

3.2 调度参数计算

real_core_num = GetCoreNumAiv()          // 获取可用 AIV 核数(如 64)
thread_num_per_block = min(2048, 32 的整数倍)
block_num = ceil(total_elements / thread_num_per_block)

// 约束检查
if block_num > 65535:
    block_num = 65535
    thread_num_per_block = ceil(total_elements / 65535)
    thread_num_per_block = ceil(thread_num_per_block / 32) * 32  // 对齐到 32

3.3 UB 内存布局

UB 总大小: 256KB
├── 静态内存(编译期确定)
├── 动态内存(dyn_ubuf_size 指定)
├── 预留空间(8KB,固定)
└── Data Cache(32KB ~ 128KB,SIMT 专用缓存)

GELU 算子不使用静态/动态内存,全部留给 Data Cache 作为访存加速。

4. Kernel 实现设计

4.1 Kernel 函数原型

template <typename Tin, typename Tout>
__global__ __launch_bounds__(2048)
void gelu_kernel(Tin* x, Tout* y, uint32_t total_elements)

4.2 核心计算逻辑

template <typename Tin, typename Tout>
__global__ __launch_bounds__(2048)
void gelu_kernel(Tin* x, Tout* y, uint32_t total_elements)
{
    // 1. 计算全局线程索引
    uint32_t idx = blockIdx.x * blockDim.x + threadIdx.x;

    if (idx >= total_elements) {
        return;
    }

    // 2. 读取输入(类型转换)
    float x_val = static_cast<float>(x[idx]);

    // 3. GELU 计算
    constexpr float COEFF_A = 0.044715f;
    constexpr float COEFF_B = -1.595769f;

    float x3 = x_val * x_val * x_val;           // x³
    float linear_part = x_val + COEFF_A * x3;    // x + 0.044715·x³
    float exp_arg = COEFF_B * linear_part;       // -1.595769·(...)
    float exp_val = expf(exp_arg);               // e^(...)
    float denom = 1.0f + exp_val;                // 1 + e^(...)
    float result = x_val / denom;                // x / (1 + e^(...))

    // 4. 写入输出(类型转换)
    y[idx] = static_cast<Tout>(result);
}

4.3 计算步骤分解

步骤 计算内容 SIMT 数学函数 说明
1 x³ = x · x · x 原生 * 立方项
2 linear = x + 0.044715 · x³ 原生 + * 线性组合
3 exp_arg = -1.595769 · linear 原生 * 系数缩放
4 exp_val = e^(exp_arg) expf() 指数函数
5 denom = 1.0 + exp_val 原生 + 分母
6 result = x / denom 原生 / 最终结果

4.4 Warp Divergence 分析

GELU 算子中所有线程执行相同的计算指令(无条件分支),不存在 Warp Divergence,硬件利用率可达 100%。


5. Host 侧实现设计

5.1 调度函数

template <typename Tin, typename Tout>
void run_gelu_dispatch(Tin* input, Tout* output, uint32_t total_elements)
{
    // 1. ACL 初始化
    aclInit(nullptr);
    int32_t deviceId = 0;
    aclrtSetDevice(deviceId);
    aclrtStream stream = nullptr;
    aclrtCreateStream(&stream);

    // 2. 内存分配
    size_t inputByteSize = total_elements * sizeof(Tin);
    size_t outputByteSize = total_elements * sizeof(Tout);

    Tin* inputHost = nullptr;
    Tout* outputHost = nullptr;
    aclrtMallocHost((void**)(&inputHost), inputByteSize);
    aclrtMallocHost((void**)(&outputHost), outputByteSize);

    Tin* inputDevice = nullptr;
    Tout* outputDevice = nullptr;
    aclrtMalloc((void**)(&inputDevice), inputByteSize, ACL_MEM_MALLOC_HUGE_FIRST);
    aclrtMalloc((void**)(&outputDevice), outputByteSize, ACL_MEM_MALLOC_HUGE_FIRST);

    // 3. Host → Device
    aclrtMemcpy(inputDevice, inputByteSize, inputHost, inputByteSize, ACL_MEMCPY_HOST_TO_DEVICE);

    // 4. 调度参数计算
    uint32_t block_num, thread_num_per_block;
    compute_launch_params(total_elements, block_num, thread_num_per_block);

    // 5. Kernel 启动
    uint32_t dyn_ubuf_size = 0;
    gelu_kernel<Tin, Tout><<<block_num, thread_num_per_block, dyn_ubuf_size, stream>>>(
        inputDevice, outputDevice, total_elements);

    // 6. 同步 + Device → Host
    aclrtSynchronizeStream(stream);
    aclrtMemcpy(outputHost, outputByteSize, outputDevice, outputByteSize, ACL_MEMCPY_DEVICE_TO_HOST);

    // 7. 资源释放
    aclrtFree(inputDevice);
    aclrtFree(outputDevice);
    aclrtFreeHost(inputHost);
    aclrtFreeHost(outputHost);
    aclrtDestroyStream(stream);
    aclrtResetDevice(deviceId);
    aclFinalize();
}

5.2 调度参数计算函数

constexpr uint32_t MAX_THREAD_COUNT = 2048;
constexpr uint32_t MAX_BLOCK_COUNT = 65535;

void compute_launch_params(uint32_t total_elements, uint32_t &block_num, uint32_t &thread_num)
{
    uint32_t real_core_num = get_core_num_aiv();  // 如 64

    // 方案1:按核数分配
    block_num = real_core_num;
    thread_num = (total_elements + block_num - 1) / block_num;

    // 对齐到 32(warp 大小)
    thread_num = ((thread_num + 31) / 32) * 32;

    if (thread_num > MAX_THREAD_COUNT) {
        thread_num = MAX_THREAD_COUNT;
        thread_num = ((thread_num + 31) / 32) * 32;  // 保持 32 对齐
        block_num = (total_elements + thread_num - 1) / thread_num;

        if (block_num > MAX_BLOCK_COUNT) {
            // 超出硬件限制
            std::cerr << "[ERROR] total_elements too large" << std::endl;
            return;
        }
    }
}

6. 工程结构设计

6.1 目录结构

gelu_simt/
├── CMakeLists.txt              # 构建配置
├── gelu_simt.asc               # SIMT kernel + host 代码
├── data_utils.h                # 文件读写工具
├── scripts/
│   ├── gen_data.py             # 输入数据和 golden 生成
│   └── verify_result.py        # 精度校验
└── README.md                   # 算子说明文档

6.2 CMakeLists.txt 配置

cmake_minimum_required(VERSION 3.16)

set(CMAKE_ASC_RUN_MODE "npu" CACHE STRING "Run mode: npu, sim")
set(CMAKE_ASC_ARCHITECTURES "dav-3510" CACHE STRING "NPU architecture: dav-3510")

find_package(ASC REQUIRED)
project(gelu_simt LANGUAGES ASC CXX)

add_executable(demo
    gelu_simt.asc
)

target_compile_options(demo PRIVATE
    $<$<COMPILE_LANGUAGE:ASC>:--npu-arch=${CMAKE_ASC_ARCHITECTURES}>
)

7. 精度验证设计

7.1 Golden 数据生成

import numpy as np

def gen_golden_data(shape=[8192, 8192]):
    input_x = np.random.uniform(-10, 10, shape).astype(np.float32)

    COEFF_A = 0.044715
    COEFF_B = -1.595769

    x3 = input_x ** 3
    linear_part = input_x + COEFF_A * x3
    exponent = COEFF_B * linear_part
    golden = input_x / (1 + np.exp(exponent))

    input_x.tofile("./input/input_x.bin")
    golden.astype(np.float32).tofile("./output/golden.bin")

7.2 精度校验

import numpy as np

RELATIVE_TOL = 1e-4
ABSOLUTE_TOL = 1e-5
ERROR_TOL = 1e-4

def verify_result(output_file, golden_file):
    output = np.fromfile(output_file, dtype=np.float32).reshape(-1)
    golden = np.fromfile(golden_file, dtype=np.float32).reshape(-1)

    different_element_results = np.isclose(output, golden,
                                           rtol=RELATIVE_TOL,
                                           atol=ABSOLUTE_TOL,
                                           equal_nan=True)
    different_element_indexes = np.where(different_element_results == False)[0]
    error_ratio = float(different_element_indexes.size) / golden.size

    print("error ratio: %.4f, tolerance: %.4f" % (error_ratio, ERROR_TOL))
    return error_ratio <= ERROR_TOL

8. 性能分析与优化

8.1 性能瓶颈分析

GELU 是纯逐元素计算,SIMT 模式下的性能瓶颈主要在于:

瓶颈类型 说明 占比预估
GM 访存带宽 每个线程读写 GM,受限于 HBM 带宽 ~60%
数学函数延迟 expf() 的硬件执行延迟 ~30%
控制流开销 blockIdx/threadIdx 计算 ~10%

8.2 优化方向

优化手段 描述 预期收益
Warp 对齐线程数 thread_num_per_block 设为 32 的整数倍 消除空闲 warp 通道
充分利用 Data Cache 不使用静态/动态内存,留出最大 Data Cache 空间 提升 GM 访存效率
增加 block_num 充分利用所有 AIV 核 提升并行度
half 精度计算 输入输出使用 half 类型,减少 GM 带宽 带宽减半,吞吐量翻倍

8.3 SIMD vs SIMT 性能对比预期

指标 SIMD (RegBase) SIMT 说明
编程复杂度 中(需理解 RegBase/VF 融合) 低(类 CUDA 风格) SIMT 更直观
向量化效率 高(一次处理 64 元素) 中(每线程 1 元素) SIMD 更适合大批量
GM 带宽利用 中(需 DataCopyPad) 高(直接 GM 访问) SIMT 有 Data Cache 加速
端到端耗时 参考基线 ~352μs 预期 ~400-500μs SIMT 略慢但差异可控
开发效率 2-3 天 0.5-1 天 SIMT 开发更快

9. 编译运行指南

9.1 编译命令

# 配置环境变量
source /usr/local/Ascend/cann-9.1.0-beta.1/set_env.sh

# 编译
mkdir -p build && cd build
cmake .. -DCMAKE_ASC_ARCHITECTURES=dav-3510 -DCMAKE_ASC_RUN_MODE=npu
make -j

# 生成测试数据
python3 ../scripts/gen_data.py

# 运行
./demo

# 精度校验
python3 ../scripts/verify_result.py output/output.bin output/golden.bin

9.2 性能分析

# 性能 profiling
msprof op ./demo

# 查看结果
cat ./OPPROF_*/OpBasicInfo.csv
cat ./OPPROF_*/PipeUtilization.csv

10. 调试工具

10.1 printf 调试

在 kernel 中使用 printf 输出调试信息:

#include "asc_printf.h"

__global__ void gelu_kernel(float* x, float* y, uint32_t total_elements)
{
    uint32_t idx = blockIdx.x * blockDim.x + threadIdx.x;
    if (idx < total_elements && idx < 10) {
        float x_val = x[idx];
        printf("thread %d: x = %f, gelu(x) = %f\n", idx, x_val, y[idx]);
    }
}

10.2 assert 调试

#include "asc_assert.h"

__global__ void gelu_kernel(float* x, float* y, uint32_t total_elements)
{
    uint32_t idx = blockIdx.x * blockDim.x + threadIdx.x;
    asc_assert(idx < total_elements, "index out of bounds");
}

11. 风险与约束

风险项 描述 应对措施
大 shape 超出硬件限制 total_elements > 2048 × 65535 在 host 侧做约束检查,超限报错
expf 数值溢出 输入值过大导致 exp 溢出 输入范围限制在 [-10, 10] 内测试
half 精度损失 half 类型精度低于 float 对 half 模式单独提高容差阈值
Data Cache 不足 静态内存分配过大导致 Data Cache < 32KB GELU 不使用静态内存,避免此风险

12. 参考文档

文档 路径
SIMT 编程简介 docs/api/SIMT-API/SIMT编程简介/
SIMT 编程模型 docs/api/SIMT-API/SIMT编程简介/编程模型.md
SIMT API 列表 docs/api/SIMT-API/SIMT编程简介/API列表.md
数学函数 docs/api/SIMT-API/数学函数/
Softmax SIMT 样例 examples/03_simt_api/00_introduction/03_softmaxv2/softmaxv2.asc
QuickStart examples/03_simt_api/00_introduction/00_quickstart/hello_world_simt/
SIMD GELU 样例 examples/01_simd_cpp_api/00_introduction/04_vector_reg/gelu/
GELU 性能调优 examples/01_simd_cpp_api/04_best_practices/02_reg_vector_compute_practices/gelu_high_performance/
Logo

鲲鹏昇腾开发者社区是面向全社会开放的“联接全球计算开发者,聚合华为+生态”的社区,内容涵盖鲲鹏、昇腾资源,帮助开发者快速获取所需的知识、经验、软件、工具、算力,支撑开发者易学、好用、成功,成为核心开发者。

更多推荐