GPUCodeForces/S1/13/evonorm_cuda.py

263 lines
9.6 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

import torch
import torch.nn as nn
from torch.utils.cpp_extension import load_inline
# 定义维度常量
N, C, H, W = 32, 64, 56, 56
EPS = 1e-6
assert (H * W) % 4 == 0, "Instance size (H * W) must be a multiple of 4"
class ModelNew(nn.Module):
"""
EvoNorm-S0/B0 的 CUDA 优化实现
"""
def __init__(self, evonorm_gamma, evonorm_beta, evonorm_v=None, use_b0=False):
super().__init__()
self.gamma = nn.Parameter(evonorm_gamma.clone().view(1, C, 1, 1))
self.beta = nn.Parameter(evonorm_beta.clone().view(1, C, 1, 1))
self.eps = EPS
self.nonlinear = (evonorm_v is not None)
self.use_b0 = use_b0
if self.nonlinear:
self.v = nn.Parameter(evonorm_v.clone().view(1, C, 1, 1))
else:
self.register_parameter('v', None)
if self.use_b0:
self.register_buffer('running_var', torch.ones(1, C, 1, 1))
self.momentum = 0.1
self._compile_cuda_kernel()
def _compile_cuda_kernel(self):
cpp_source = """
#include <torch/extension.h>
torch::Tensor evonorm_forward_cuda(
torch::Tensor input,
torch::Tensor mean,
torch::Tensor var,
torch::Tensor gamma,
torch::Tensor beta,
torch::Tensor v,
float eps,
bool nonlinear,
int N, int C, int H, int W);
"""
cuda_source = """
#include <cuda_runtime.h>
#include <device_launch_parameters.h>
#include <cmath>
// 优化: 使用快速数学函数
#define FAST_DIV(a, b) __fdividef(a, b)
#define FAST_EXP(x) __expf(x)
// 优化 1: Sigmoid 快速计算(使用查找表或优化公式)
__device__ __forceinline__ float fast_sigmoid(float x) {
// 使用快速除法和指数
return FAST_DIV(1.0f, 1.0f + FAST_EXP(-x));
}
// 优化 2: 向量化 sigmoid 计算
__device__ __forceinline__ float4 sigmoid_vec(float4 x, float v_val) {
float4 result;
result.x = fast_sigmoid(x.x * v_val);
result.y = fast_sigmoid(x.y * v_val);
result.z = fast_sigmoid(x.z * v_val);
result.w = fast_sigmoid(x.w * v_val);
return result;
}
__global__ void evonorm_apply_kernel(
const float* __restrict__ x,
const float* __restrict__ mean,
const float* __restrict__ var,
const float* __restrict__ gamma,
const float* __restrict__ beta,
const float* __restrict__ v,
float* __restrict__ y,
float eps,
bool nonlinear,
int N, int C, int H, int W
) {
const int nc_idx = blockIdx.x;
if (nc_idx >= N * C) return;
const int n_idx = nc_idx / C;
const int c_idx = nc_idx % C;
// 优化 3: 使用 __ldg() 读取只读全局内存
const float m = __ldg(&mean[nc_idx]);
const float variance = __ldg(&var[nc_idx]);
// 优化 4: 预计算常量
const float inv_std = rsqrtf(variance + eps); // rsqrtf 比 1.0f/sqrtf 快
const float g = __ldg(&gamma[c_idx]);
const float b = __ldg(&beta[c_idx]);
const float v_val = nonlinear ? __ldg(&v[c_idx]) : 0.0f;
const int instance_size = H * W;
const int instance_offset = n_idx * C * instance_size + c_idx * instance_size;
const float* x_ptr = x + instance_offset;
float* y_ptr = y + instance_offset;
const int instance_size_div4 = instance_size / 4;
const float4* x4_ptr = reinterpret_cast<const float4*>(x_ptr);
float4* y4_ptr = reinterpret_cast<float4*>(y_ptr);
const int BLOCK_SIZE = 256;
// 优化 5: 循环展开(处理 2 个 float4 每次迭代)
const int items_per_thread = (instance_size_div4 + BLOCK_SIZE - 1) / BLOCK_SIZE;
const int base_idx = threadIdx.x;
#pragma unroll 2
for (int i = 0; i < items_per_thread; ++i) {
int idx = base_idx + i * BLOCK_SIZE;
if (idx < instance_size_div4) {
// 优化 6: 使用 __ldg() 读取输入(如果对齐)
float4 x_val = x4_ptr[idx];
float4 y_val;
// 归一化: (x - m) / std
// 注意: 不使用 volatile因为统计量已在 Python 端计算
float x_norm_x = (x_val.x - m) * inv_std;
float x_norm_y = (x_val.y - m) * inv_std;
float x_norm_z = (x_val.z - m) * inv_std;
float x_norm_w = (x_val.w - m) * inv_std;
// 仿射变换: x_norm * g + b (使用 FMA)
float y_affine_x = fmaf(x_norm_x, g, b);
float y_affine_y = fmaf(x_norm_y, g, b);
float y_affine_z = fmaf(x_norm_z, g, b);
float y_affine_w = fmaf(x_norm_w, g, b);
// 非线性门控
if (nonlinear) {
// 优化 7: 向量化 sigmoid 计算
float sigmoid_x = fast_sigmoid(x_val.x * v_val);
float sigmoid_y = fast_sigmoid(x_val.y * v_val);
float sigmoid_z = fast_sigmoid(x_val.z * v_val);
float sigmoid_w = fast_sigmoid(x_val.w * v_val);
y_val.x = y_affine_x * sigmoid_x;
y_val.y = y_affine_y * sigmoid_y;
y_val.z = y_affine_z * sigmoid_z;
y_val.w = y_affine_w * sigmoid_w;
} else {
y_val.x = y_affine_x;
y_val.y = y_affine_y;
y_val.z = y_affine_z;
y_val.w = y_affine_w;
}
y4_ptr[idx] = y_val;
}
}
}
// ============================================================
// C++ Wrapper
// ============================================================
torch::Tensor evonorm_forward_cuda(
torch::Tensor input,
torch::Tensor mean,
torch::Tensor var,
torch::Tensor gamma,
torch::Tensor beta,
torch::Tensor v,
float eps,
bool nonlinear,
int N, int C, int H, int W
) {
input = input.contiguous();
auto output = torch::empty_like(input);
const int BLOCK_SIZE = 256;
dim3 blocks(N * C);
dim3 threads(BLOCK_SIZE);
// 优化 8: 使用 CUDA stream可选
evonorm_apply_kernel<<<blocks, threads>>>(
input.data_ptr<float>(),
mean.data_ptr<float>(),
var.data_ptr<float>(),
gamma.data_ptr<float>(),
beta.data_ptr<float>(),
nonlinear ? v.data_ptr<float>() : nullptr,
output.data_ptr<float>(),
eps,
nonlinear,
N, C, H, W
);
return output;
}
"""
# 优化 9: 使用更激进的编译选项
self.evonorm_op = load_inline(
name="evonorm_cuda_optimized_v4",
cpp_sources=cpp_source,
cuda_sources=cuda_source,
functions=["evonorm_forward_cuda"],
extra_cuda_cflags=[
"-O3",
"--use_fast_math", # 启用快速数学(可能略微降低精度但提升性能)
"-lineinfo" # 便于性能分析
],
verbose=False
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
if x.dtype != torch.float32 or not x.is_cuda:
x = x.to("cuda", dtype=torch.float32)
N, C, H, W = x.size()
if self.use_b0:
# EvoNorm-B0
if self.training:
mean = x.mean(dim=[2, 3], keepdim=True)
var = x.var(dim=[2, 3], keepdim=True, unbiased=False)
with torch.no_grad():
batch_var = var.mean(dim=0, keepdim=True)
self.running_var = (1 - self.momentum) * self.running_var + self.momentum * batch_var
else:
mean = x.mean(dim=[2, 3], keepdim=True)
var = self.running_var.expand(N, C, 1, 1)
else:
# EvoNorm-S0
# 优化 10: 融合计算 E[x^2] 和 E[x] 可以考虑自定义 CUDA kernel
x_sq_mean = torch.mean(x * x, dim=[2, 3], keepdim=True)
x_mean = torch.mean(x, dim=[2, 3], keepdim=True)
var = x_sq_mean - x_mean * x_mean
mean = torch.zeros_like(x_mean)
gamma_view = self.gamma.data.view(C).contiguous()
beta_view = self.beta.data.view(C).contiguous()
if self.nonlinear:
v_view = self.v.data.view(C).contiguous()
else:
v_view = torch.zeros(C, device=x.device, dtype=torch.float32)
return self.evonorm_op.evonorm_forward_cuda(
x.contiguous(),
mean.contiguous().view(N, C),
var.contiguous().view(N, C),
gamma_view,
beta_view,
v_view,
self.eps,
self.nonlinear,
N, C, H, W
)