2025-10-29 15:09:28 +08:00
|
|
|
|
import torch
|
|
|
|
|
|
import torch.nn as nn
|
|
|
|
|
|
from torch.utils.cpp_extension import load_inline
|
|
|
|
|
|
|
|
|
|
|
|
batchnorm_source = r"""
|
|
|
|
|
|
#include <torch/extension.h>
|
|
|
|
|
|
#include <cuda_runtime.h>
|
|
|
|
|
|
#include <ATen/cuda/CUDAContext.h>
|
|
|
|
|
|
#include <c10/cuda/CUDAException.h>
|
|
|
|
|
|
|
2025-10-30 12:32:48 +08:00
|
|
|
|
// 统一的训练 kernel(计算批次统计量)
|
|
|
|
|
|
__global__ void batchnorm_forward_train_kernel_optimized(
|
2025-10-29 15:09:28 +08:00
|
|
|
|
const float* __restrict__ x,
|
|
|
|
|
|
const float* __restrict__ gamma,
|
|
|
|
|
|
const float* __restrict__ beta,
|
2025-10-30 12:32:48 +08:00
|
|
|
|
float* __restrict__ running_mean,
|
|
|
|
|
|
float* __restrict__ running_var,
|
2025-10-29 15:09:28 +08:00
|
|
|
|
float* __restrict__ y,
|
|
|
|
|
|
int batch,
|
|
|
|
|
|
int features,
|
2025-10-30 12:32:48 +08:00
|
|
|
|
float eps,
|
|
|
|
|
|
float momentum,
|
|
|
|
|
|
bool update_stats // 是否更新统计量
|
2025-10-29 15:09:28 +08:00
|
|
|
|
) {
|
|
|
|
|
|
int feature = blockIdx.x;
|
|
|
|
|
|
if (feature >= features) return;
|
2025-10-30 12:32:48 +08:00
|
|
|
|
|
2025-10-29 15:09:28 +08:00
|
|
|
|
int tid = threadIdx.x;
|
2025-10-30 12:32:48 +08:00
|
|
|
|
int num_threads = blockDim.x;
|
|
|
|
|
|
int warp_id = tid / 32;
|
|
|
|
|
|
int lane_id = tid % 32;
|
|
|
|
|
|
int num_warps = (num_threads + 31) / 32;
|
|
|
|
|
|
|
|
|
|
|
|
const float* x_base = x + feature;
|
|
|
|
|
|
float* y_base = y + feature;
|
|
|
|
|
|
|
2025-10-29 15:09:28 +08:00
|
|
|
|
float sum = 0.0f;
|
|
|
|
|
|
float sum_sq = 0.0f;
|
2025-10-30 12:32:48 +08:00
|
|
|
|
|
|
|
|
|
|
int row = tid;
|
|
|
|
|
|
for (; row + num_threads <= batch; row += num_threads) {
|
|
|
|
|
|
float v = x_base[row * features];
|
2025-10-29 15:09:28 +08:00
|
|
|
|
sum += v;
|
|
|
|
|
|
sum_sq += v * v;
|
|
|
|
|
|
}
|
2025-10-30 12:32:48 +08:00
|
|
|
|
|
|
|
|
|
|
if (row < batch) {
|
|
|
|
|
|
float v = x_base[row * features];
|
|
|
|
|
|
sum += v;
|
|
|
|
|
|
sum_sq += v * v;
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
#pragma unroll
|
|
|
|
|
|
for (int offset = 16; offset > 0; offset >>= 1) {
|
|
|
|
|
|
sum += __shfl_down_sync(0xffffffff, sum, offset);
|
|
|
|
|
|
sum_sq += __shfl_down_sync(0xffffffff, sum_sq, offset);
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
__shared__ float shared_sum[32];
|
|
|
|
|
|
__shared__ float shared_sq[32];
|
|
|
|
|
|
|
|
|
|
|
|
if (lane_id == 0) {
|
|
|
|
|
|
shared_sum[warp_id] = sum;
|
|
|
|
|
|
shared_sq[warp_id] = sum_sq;
|
|
|
|
|
|
}
|
2025-10-29 15:09:28 +08:00
|
|
|
|
__syncthreads();
|
2025-10-30 12:32:48 +08:00
|
|
|
|
|
|
|
|
|
|
if (tid < 32) {
|
|
|
|
|
|
sum = (tid < num_warps) ? shared_sum[tid] : 0.0f;
|
|
|
|
|
|
sum_sq = (tid < num_warps) ? shared_sq[tid] : 0.0f;
|
|
|
|
|
|
|
|
|
|
|
|
#pragma unroll
|
|
|
|
|
|
for (int offset = 16; offset > 0; offset >>= 1) {
|
|
|
|
|
|
sum += __shfl_down_sync(0xffffffff, sum, offset);
|
|
|
|
|
|
sum_sq += __shfl_down_sync(0xffffffff, sum_sq, offset);
|
2025-10-29 15:09:28 +08:00
|
|
|
|
}
|
|
|
|
|
|
}
|
2025-10-30 12:32:48 +08:00
|
|
|
|
|
2025-10-29 15:09:28 +08:00
|
|
|
|
__shared__ float s_mean;
|
|
|
|
|
|
__shared__ float s_inv_std;
|
2025-10-30 12:32:48 +08:00
|
|
|
|
__shared__ float s_gamma;
|
|
|
|
|
|
__shared__ float s_beta;
|
|
|
|
|
|
|
2025-10-29 15:09:28 +08:00
|
|
|
|
if (tid == 0) {
|
2025-10-30 12:32:48 +08:00
|
|
|
|
float mean = sum / batch;
|
|
|
|
|
|
float var = (sum_sq / batch) - (mean * mean);
|
|
|
|
|
|
var = fmaxf(var, 0.0f);
|
2025-10-29 15:09:28 +08:00
|
|
|
|
s_mean = mean;
|
|
|
|
|
|
s_inv_std = rsqrtf(var + eps);
|
2025-10-30 12:32:48 +08:00
|
|
|
|
s_gamma = gamma[feature];
|
|
|
|
|
|
s_beta = beta[feature];
|
|
|
|
|
|
|
|
|
|
|
|
// 只有需要时才更新 running stats
|
|
|
|
|
|
if (update_stats) {
|
|
|
|
|
|
running_mean[feature] = (1.0f - momentum) * running_mean[feature] + momentum * mean;
|
|
|
|
|
|
float unbiased_var = var * batch / fmaxf(float(batch - 1), 1.0f);
|
|
|
|
|
|
running_var[feature] = (1.0f - momentum) * running_var[feature] + momentum * unbiased_var;
|
|
|
|
|
|
}
|
2025-10-29 15:09:28 +08:00
|
|
|
|
}
|
|
|
|
|
|
__syncthreads();
|
2025-10-30 12:32:48 +08:00
|
|
|
|
|
2025-10-29 15:09:28 +08:00
|
|
|
|
float mean = s_mean;
|
|
|
|
|
|
float inv_std = s_inv_std;
|
2025-10-30 12:32:48 +08:00
|
|
|
|
float g = s_gamma;
|
|
|
|
|
|
float b = s_beta;
|
|
|
|
|
|
|
|
|
|
|
|
row = tid;
|
|
|
|
|
|
for (; row + num_threads <= batch; row += num_threads) {
|
|
|
|
|
|
float v = x_base[row * features];
|
|
|
|
|
|
float norm = (v - mean) * inv_std;
|
|
|
|
|
|
y_base[row * features] = norm * g + b;
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
if (row < batch) {
|
|
|
|
|
|
float v = x_base[row * features];
|
|
|
|
|
|
float norm = (v - mean) * inv_std;
|
|
|
|
|
|
y_base[row * features] = norm * g + b;
|
|
|
|
|
|
}
|
|
|
|
|
|
}
|
2025-10-29 15:09:28 +08:00
|
|
|
|
|
2025-10-30 12:32:48 +08:00
|
|
|
|
// 推理模式 kernel(使用 running stats)
|
|
|
|
|
|
__global__ void batchnorm_forward_eval_kernel_optimized(
|
|
|
|
|
|
const float* __restrict__ x,
|
|
|
|
|
|
const float* __restrict__ gamma,
|
|
|
|
|
|
const float* __restrict__ beta,
|
|
|
|
|
|
const float* __restrict__ running_mean,
|
|
|
|
|
|
const float* __restrict__ running_var,
|
|
|
|
|
|
float* __restrict__ y,
|
|
|
|
|
|
int batch,
|
|
|
|
|
|
int features,
|
|
|
|
|
|
float eps
|
|
|
|
|
|
) {
|
|
|
|
|
|
int tid = blockIdx.x * blockDim.x + threadIdx.x;
|
|
|
|
|
|
int total = batch * features;
|
|
|
|
|
|
int stride = gridDim.x * blockDim.x;
|
|
|
|
|
|
|
|
|
|
|
|
for (int idx = tid; idx < total; idx += stride) {
|
|
|
|
|
|
int feature = idx % features;
|
|
|
|
|
|
|
|
|
|
|
|
float mean = running_mean[feature];
|
|
|
|
|
|
float var = running_var[feature];
|
|
|
|
|
|
float inv_std = rsqrtf(var + eps);
|
|
|
|
|
|
float g = gamma[feature];
|
|
|
|
|
|
float b = beta[feature];
|
|
|
|
|
|
|
|
|
|
|
|
float v = x[idx];
|
2025-10-29 15:09:28 +08:00
|
|
|
|
float norm = (v - mean) * inv_std;
|
2025-10-30 12:32:48 +08:00
|
|
|
|
y[idx] = norm * g + b;
|
2025-10-29 15:09:28 +08:00
|
|
|
|
}
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
torch::Tensor batchnorm_cuda_forward(
|
|
|
|
|
|
torch::Tensor x,
|
|
|
|
|
|
torch::Tensor weight,
|
|
|
|
|
|
torch::Tensor bias,
|
2025-10-30 12:32:48 +08:00
|
|
|
|
torch::Tensor running_mean,
|
|
|
|
|
|
torch::Tensor running_var,
|
|
|
|
|
|
bool training,
|
|
|
|
|
|
double momentum,
|
|
|
|
|
|
double eps,
|
|
|
|
|
|
bool track_running_stats // 改名:更清晰地表达意图
|
2025-10-29 15:09:28 +08:00
|
|
|
|
) {
|
|
|
|
|
|
TORCH_CHECK(x.is_cuda(), "x must be a CUDA tensor");
|
|
|
|
|
|
TORCH_CHECK(weight.is_cuda(), "weight must be a CUDA tensor");
|
|
|
|
|
|
TORCH_CHECK(bias.is_cuda(), "bias must be a CUDA tensor");
|
|
|
|
|
|
TORCH_CHECK(x.dtype() == torch::kFloat32, "only float32 tensors are supported");
|
|
|
|
|
|
TORCH_CHECK(weight.dtype() == torch::kFloat32, "weight must be float32");
|
|
|
|
|
|
TORCH_CHECK(bias.dtype() == torch::kFloat32, "bias must be float32");
|
|
|
|
|
|
TORCH_CHECK(x.dim() == 2, "input must be 2D [batch, features]");
|
|
|
|
|
|
TORCH_CHECK(weight.dim() == 1, "weight must be 1D");
|
|
|
|
|
|
TORCH_CHECK(bias.dim() == 1, "bias must be 1D");
|
2025-10-30 12:32:48 +08:00
|
|
|
|
TORCH_CHECK(x.size(1) == weight.size(0), "feature size mismatch");
|
2025-10-29 15:09:28 +08:00
|
|
|
|
TORCH_CHECK(weight.size(0) == bias.size(0), "weight and bias must have the same length");
|
|
|
|
|
|
|
|
|
|
|
|
auto x_contig = x.contiguous();
|
|
|
|
|
|
auto weight_contig = weight.contiguous();
|
|
|
|
|
|
auto bias_contig = bias.contiguous();
|
|
|
|
|
|
|
|
|
|
|
|
int batch = x_contig.size(0);
|
|
|
|
|
|
int features = x_contig.size(1);
|
|
|
|
|
|
|
|
|
|
|
|
auto y = torch::empty_like(x_contig);
|
2025-10-30 12:32:48 +08:00
|
|
|
|
cudaStream_t stream = at::cuda::getCurrentCUDAStream();
|
2025-10-29 15:09:28 +08:00
|
|
|
|
|
2025-10-30 12:32:48 +08:00
|
|
|
|
TORCH_CHECK(running_mean.is_cuda(), "running_mean must be a CUDA tensor");
|
|
|
|
|
|
TORCH_CHECK(running_var.is_cuda(), "running_var must be a CUDA tensor");
|
|
|
|
|
|
TORCH_CHECK(running_mean.dim() == 1, "running_mean must be 1D");
|
|
|
|
|
|
TORCH_CHECK(running_var.dim() == 1, "running_var must be 1D");
|
|
|
|
|
|
TORCH_CHECK(running_mean.size(0) == features, "running_mean size mismatch");
|
|
|
|
|
|
TORCH_CHECK(running_var.size(0) == features, "running_var size mismatch");
|
2025-10-29 15:09:28 +08:00
|
|
|
|
|
2025-10-30 12:32:48 +08:00
|
|
|
|
// 关键修改:根据 track_running_stats 决定行为
|
|
|
|
|
|
// track_running_stats=False: 总是计算批次统计(训练和推理都一样)
|
|
|
|
|
|
// track_running_stats=True + training: 计算批次统计并更新 running stats
|
|
|
|
|
|
// track_running_stats=True + eval: 使用 running stats
|
|
|
|
|
|
|
|
|
|
|
|
bool use_batch_stats = !track_running_stats || training;
|
|
|
|
|
|
|
|
|
|
|
|
if (use_batch_stats) {
|
|
|
|
|
|
// 使用批次统计量(训练模式 或 track_running_stats=False)
|
|
|
|
|
|
int threads;
|
|
|
|
|
|
if (batch <= 16) {
|
|
|
|
|
|
threads = 32;
|
|
|
|
|
|
} else if (batch <= 32) {
|
|
|
|
|
|
threads = 32;
|
|
|
|
|
|
} else if (batch <= 64) {
|
|
|
|
|
|
threads = 64;
|
|
|
|
|
|
} else if (batch <= 128) {
|
|
|
|
|
|
threads = 128;
|
|
|
|
|
|
} else if (batch <= 256) {
|
|
|
|
|
|
threads = 256;
|
|
|
|
|
|
} else {
|
|
|
|
|
|
threads = 256;
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
int blocks = features;
|
|
|
|
|
|
size_t shared_mem = 0;
|
|
|
|
|
|
|
|
|
|
|
|
// update_stats = track_running_stats && training
|
|
|
|
|
|
// track_running_stats=False: 不更新
|
|
|
|
|
|
// track_running_stats=True + training: 更新
|
|
|
|
|
|
// track_running_stats=True + eval: 不会走到这里
|
|
|
|
|
|
bool update_stats = track_running_stats && training;
|
|
|
|
|
|
|
|
|
|
|
|
batchnorm_forward_train_kernel_optimized<<<blocks, threads, shared_mem, stream>>>(
|
|
|
|
|
|
x_contig.data_ptr<float>(),
|
|
|
|
|
|
weight_contig.data_ptr<float>(),
|
|
|
|
|
|
bias_contig.data_ptr<float>(),
|
|
|
|
|
|
running_mean.data_ptr<float>(),
|
|
|
|
|
|
running_var.data_ptr<float>(),
|
|
|
|
|
|
y.data_ptr<float>(),
|
|
|
|
|
|
batch,
|
|
|
|
|
|
features,
|
|
|
|
|
|
static_cast<float>(eps),
|
|
|
|
|
|
static_cast<float>(momentum),
|
|
|
|
|
|
update_stats
|
|
|
|
|
|
);
|
|
|
|
|
|
} else {
|
|
|
|
|
|
// 使用 running stats(track_running_stats=True + eval 模式)
|
|
|
|
|
|
int total = batch * features;
|
|
|
|
|
|
int threads = 256;
|
|
|
|
|
|
int blocks;
|
|
|
|
|
|
|
|
|
|
|
|
if (total <= 4096) {
|
|
|
|
|
|
blocks = (total + threads - 1) / threads;
|
|
|
|
|
|
} else {
|
|
|
|
|
|
blocks = min(1024, (total + threads * 4 - 1) / (threads * 4));
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
batchnorm_forward_eval_kernel_optimized<<<blocks, threads, 0, stream>>>(
|
|
|
|
|
|
x_contig.data_ptr<float>(),
|
|
|
|
|
|
weight_contig.data_ptr<float>(),
|
|
|
|
|
|
bias_contig.data_ptr<float>(),
|
|
|
|
|
|
running_mean.data_ptr<float>(),
|
|
|
|
|
|
running_var.data_ptr<float>(),
|
|
|
|
|
|
y.data_ptr<float>(),
|
|
|
|
|
|
batch,
|
|
|
|
|
|
features,
|
|
|
|
|
|
static_cast<float>(eps)
|
|
|
|
|
|
);
|
|
|
|
|
|
}
|
|
|
|
|
|
|
2025-10-29 15:09:28 +08:00
|
|
|
|
C10_CUDA_KERNEL_LAUNCH_CHECK();
|
|
|
|
|
|
return y;
|
|
|
|
|
|
}
|
|
|
|
|
|
"""
|
|
|
|
|
|
|
|
|
|
|
|
batchnorm_cpp_source = r"""
|
|
|
|
|
|
torch::Tensor batchnorm_cuda_forward(
|
|
|
|
|
|
torch::Tensor x,
|
|
|
|
|
|
torch::Tensor weight,
|
|
|
|
|
|
torch::Tensor bias,
|
2025-10-30 12:32:48 +08:00
|
|
|
|
torch::Tensor running_mean,
|
|
|
|
|
|
torch::Tensor running_var,
|
|
|
|
|
|
bool training,
|
|
|
|
|
|
double momentum,
|
|
|
|
|
|
double eps,
|
|
|
|
|
|
bool track_running_stats
|
2025-10-29 15:09:28 +08:00
|
|
|
|
);
|
|
|
|
|
|
"""
|
|
|
|
|
|
|
|
|
|
|
|
batchnorm_cuda = load_inline(
|
|
|
|
|
|
name="batchnorm_cuda_ext",
|
|
|
|
|
|
cpp_sources=batchnorm_cpp_source,
|
|
|
|
|
|
cuda_sources=batchnorm_source,
|
|
|
|
|
|
functions=["batchnorm_cuda_forward"],
|
|
|
|
|
|
verbose=True
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
class ModelNew(nn.Module):
|
|
|
|
|
|
"""
|
|
|
|
|
|
Model performing matrix multiplication followed by custom CUDA BatchNorm and ReLU.
|
2025-10-30 12:32:48 +08:00
|
|
|
|
Optimized with Warp-level reduction (Plan 1) and thread configuration (Plan 2).
|
2025-10-29 15:09:28 +08:00
|
|
|
|
"""
|
2025-10-30 12:32:48 +08:00
|
|
|
|
def __init__(self, mat_weight: torch.Tensor, bn_weight: torch.Tensor, bn_bias: torch.Tensor,
|
|
|
|
|
|
eps: float = 1e-5, momentum: float = 0.1, track_running_stats: bool = True):
|
2025-10-29 15:09:28 +08:00
|
|
|
|
super().__init__()
|
|
|
|
|
|
if mat_weight.dim() != 2:
|
|
|
|
|
|
raise ValueError("mat_weight must be a 2D tensor [input_dim, output_dim].")
|
|
|
|
|
|
if bn_weight.dim() != 1 or bn_bias.dim() != 1:
|
|
|
|
|
|
raise ValueError("BatchNorm weight and bias must be 1D.")
|
|
|
|
|
|
if bn_weight.size(0) != mat_weight.size(1):
|
|
|
|
|
|
raise ValueError("BatchNorm parameter size must match output_dim.")
|
|
|
|
|
|
if bn_weight.size(0) != bn_bias.size(0):
|
|
|
|
|
|
raise ValueError("BatchNorm weight and bias must share shape.")
|
2025-10-30 12:32:48 +08:00
|
|
|
|
|
2025-10-29 15:09:28 +08:00
|
|
|
|
self.weight = nn.Parameter(mat_weight.clone())
|
|
|
|
|
|
self.bn_weight = nn.Parameter(bn_weight.clone())
|
|
|
|
|
|
self.bn_bias = nn.Parameter(bn_bias.clone())
|
|
|
|
|
|
self.eps = eps
|
2025-10-30 12:32:48 +08:00
|
|
|
|
self.momentum = momentum
|
|
|
|
|
|
self.track_running_stats = track_running_stats
|
|
|
|
|
|
|
|
|
|
|
|
# 无论 track_running_stats 是什么,都创建 buffer
|
|
|
|
|
|
self.register_buffer('running_mean', torch.zeros(bn_weight.size(0)))
|
|
|
|
|
|
self.register_buffer('running_var', torch.ones(bn_weight.size(0)))
|
2025-10-29 15:09:28 +08:00
|
|
|
|
|
|
|
|
|
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
|
|
|
|
|
if not x.is_cuda:
|
|
|
|
|
|
raise ValueError("Input must be a CUDA tensor.")
|
|
|
|
|
|
if not self.weight.is_cuda:
|
|
|
|
|
|
raise ValueError("Model weight must be on CUDA.")
|
|
|
|
|
|
if not self.bn_weight.is_cuda or not self.bn_bias.is_cuda:
|
|
|
|
|
|
raise ValueError("BatchNorm parameters must be on CUDA.")
|
2025-10-30 12:32:48 +08:00
|
|
|
|
|
2025-10-29 15:09:28 +08:00
|
|
|
|
x = torch.matmul(x, self.weight)
|
2025-10-30 12:32:48 +08:00
|
|
|
|
|
|
|
|
|
|
# 传递 track_running_stats 参数到 CUDA kernel
|
|
|
|
|
|
x = batchnorm_cuda.batchnorm_cuda_forward(
|
|
|
|
|
|
x,
|
|
|
|
|
|
self.bn_weight,
|
|
|
|
|
|
self.bn_bias,
|
|
|
|
|
|
self.running_mean,
|
|
|
|
|
|
self.running_var,
|
|
|
|
|
|
self.training,
|
|
|
|
|
|
self.momentum,
|
|
|
|
|
|
self.eps,
|
|
|
|
|
|
self.track_running_stats
|
|
|
|
|
|
)
|
|
|
|
|
|
|
2025-10-29 15:09:28 +08:00
|
|
|
|
return torch.relu(x)
|