Softmax
2026/6/6大约 1 分钟
Softmax
原始题目:LeetGPU - Softmax
题目描述
编写一个 GPU 程序,计算 32 位浮点数数组的 softmax 函数。对于长度为 的输入数组 ,softmax 是一个长度相同的数组,其第 个元素定义为:
你的解法应使用 "max trick" 来处理潜在的溢出问题:在指数运算之前,将输入数组的每个元素减去最大值。
实现要求
- 只允许使用原生功能(不允许使用外部库)。
solve函数签名必须保持不变。- 最终结果必须存储在数组
output中。
示例
示例 1
Input: [1.0, 2.0, 3.0], N = 3
Output: [0.090, 0.244, 0.665](近似值)示例 2
Input: [-10.0, -5.0, 0.0, 5.0, 10.0], N = 5
Output: [2.047e-09, 3.038e-07, 4.509e-05, 6.693e-03, 9.933e-01](近似值)约束条件
- 。
- 性能测试在 的规模下进行。
解题思路
Softmax 需要三趟遍历:找最大值(reduce)、指数求和(reduce)、逐元素除法(map)。在线程协作层面,可以先并行找局部最大值,通过 warp shuffle 或 shared memory 规约得到全局最大值。第二轮同样的规约模式求指数和。两趟 reduce 都可以使用高效的分块规约策略。注意处理 的数值稳定性。
代码实现
CUDA
#include <cuda_runtime.h>
#include <math.h>
__global__ void softmax_kernel(const float* input, float* output, int N) {
__shared__ float sdata[256];
int tid = threadIdx.x;
// Find max
float mx = -INFINITY;
for (int i = blockIdx.x * blockDim.x + tid; i < N; i += gridDim.x * blockDim.x)
mx = fmaxf(mx, input[i]);
sdata[tid] = mx; __syncthreads();
for (int s = blockDim.x/2; s > 0; s >>= 1) {
if (tid < s) sdata[tid] = fmaxf(sdata[tid], sdata[tid+s]); __syncthreads();
}
mx = sdata[0];
// Sum exp
float sm = 0.0f;
for (int i = blockIdx.x * blockDim.x + tid; i < N; i += gridDim.x * blockDim.x)
sm += expf(input[i] - mx);
sdata[tid] = sm; __syncthreads();
for (int s = blockDim.x/2; s > 0; s >>= 1) {
if (tid < s) sdata[tid] += sdata[tid+s]; __syncthreads();
}
sm = sdata[0];
// Normalize
for (int i = blockIdx.x * blockDim.x + tid; i < N; i += gridDim.x * blockDim.x)
output[i] = expf(input[i] - mx) / sm;
}
extern "C" void solve(const float* input, float* output, int N) {
softmax_kernel<<<1, 256>>>(input, output, N);
cudaDeviceSynchronize();
}Triton
import triton, triton.language as tl
@triton.jit
def softmax_kernel(input_ptr, output_ptr, N: tl.constexpr, BLOCK: tl.constexpr):
idx = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK)
mask = idx < N
x = tl.load(input_ptr + idx, mask=mask, other=-float('inf'))
x_max = tl.max(x, axis=0)
num = tl.exp(x - x_max)
denom = tl.sum(num, axis=0)
tl.store(output_ptr + idx, num / denom, mask=mask)