forked from ccf-ai-infra/GPUCodeForces
298 lines
12 KiB
Python
298 lines
12 KiB
Python
# circleloss_cuda.py
|
||
import torch
|
||
import torch.nn.functional as F
|
||
from torch.utils.cpp_extension import load_inline
|
||
# 修复:从正确的文件导入
|
||
from circleloss_torch import BATCH_SIZE, FEATURE_DIM, MARGIN, GAMMA
|
||
|
||
|
||
class ModelNew(torch.nn.Module):
|
||
|
||
def __init__(self):
|
||
super().__init__()
|
||
self.margin = MARGIN
|
||
self.gamma = GAMMA
|
||
self._compile_cuda_kernel()
|
||
|
||
def _compile_cuda_kernel(self):
|
||
cpp_source = """
|
||
#include <torch/extension.h>
|
||
|
||
// C++ 接口 (保持不变)
|
||
torch::Tensor circleloss_forward_cuda(
|
||
torch::Tensor similarities,
|
||
torch::Tensor labels,
|
||
float margin_val,
|
||
float gamma_val
|
||
);
|
||
"""
|
||
|
||
cuda_source = """
|
||
#include <torch/extension.h>
|
||
#include <cuda_runtime.h>
|
||
#include <cmath>
|
||
#include <stdint.h> // for int64_t
|
||
#include <float.h> // For FLT_MAX
|
||
|
||
#define BLOCK_SIZE 256
|
||
|
||
// ------------------------------------------------------------------
|
||
// 阶段 1: 寻找 Logits 的最大值
|
||
// ------------------------------------------------------------------
|
||
__global__ void circleloss_find_max_kernel(
|
||
const float* __restrict__ similarities_data,
|
||
const int64_t* __restrict__ labels_data,
|
||
float* __restrict__ block_max_p_out, // (grid_size,)
|
||
float* __restrict__ block_max_n_out, // (grid_size,)
|
||
int n_elements,
|
||
int batch_size,
|
||
float margin_val,
|
||
float gamma_val
|
||
) {
|
||
__shared__ float s_data_p[BLOCK_SIZE];
|
||
__shared__ float s_data_n[BLOCK_SIZE];
|
||
|
||
float thread_max_p = -FLT_MAX;
|
||
float thread_max_n = -FLT_MAX;
|
||
|
||
const float delta_p = 1.0f - margin_val;
|
||
const float delta_n = margin_val;
|
||
|
||
int grid_stride = gridDim.x * blockDim.x;
|
||
|
||
for (int idx = blockIdx.x * blockDim.x + threadIdx.x;
|
||
idx < n_elements;
|
||
idx += grid_stride)
|
||
{
|
||
int i = idx / batch_size;
|
||
int j = idx % batch_size;
|
||
float s = similarities_data[idx];
|
||
|
||
if (labels_data[i] == labels_data[j]) {
|
||
// 正样本对
|
||
float ap = fmaxf(0.0f, -s + 1.0f + margin_val);
|
||
float logit_p = -ap * (s - delta_p) * gamma_val;
|
||
thread_max_p = fmaxf(thread_max_p, logit_p);
|
||
} else {
|
||
// 负样本对
|
||
float an = fmaxf(0.0f, s + margin_val);
|
||
float logit_n = an * (s - delta_n) * gamma_val;
|
||
thread_max_n = fmaxf(thread_max_n, logit_n);
|
||
}
|
||
}
|
||
|
||
// --- 块内归约 (Max) - 正样本对 ---
|
||
s_data_p[threadIdx.x] = thread_max_p;
|
||
__syncthreads();
|
||
for (int offset = BLOCK_SIZE / 2; offset > 0; offset >>= 1) {
|
||
if (threadIdx.x < offset) {
|
||
s_data_p[threadIdx.x] = fmaxf(s_data_p[threadIdx.x], s_data_p[threadIdx.x + offset]);
|
||
}
|
||
__syncthreads();
|
||
}
|
||
|
||
// --- 块内归约 (Max) - 负样本对 ---
|
||
s_data_n[threadIdx.x] = thread_max_n;
|
||
__syncthreads();
|
||
for (int offset = BLOCK_SIZE / 2; offset > 0; offset >>= 1) {
|
||
if (threadIdx.x < offset) {
|
||
s_data_n[threadIdx.x] = fmaxf(s_data_n[threadIdx.x], s_data_n[threadIdx.x + offset]);
|
||
}
|
||
__syncthreads();
|
||
}
|
||
|
||
if (threadIdx.x == 0) {
|
||
block_max_p_out[blockIdx.x] = s_data_p[0];
|
||
block_max_n_out[blockIdx.x] = s_data_n[0];
|
||
}
|
||
}
|
||
|
||
|
||
// ------------------------------------------------------------------
|
||
// 阶段 2: 计算 Sum(Exp(Logit - Max))
|
||
// ------------------------------------------------------------------
|
||
__global__ void circleloss_sum_exp_diff_kernel(
|
||
const float* __restrict__ similarities_data,
|
||
const int64_t* __restrict__ labels_data,
|
||
float* __restrict__ block_sum_p_out, // (grid_size,)
|
||
float* __restrict__ block_sum_n_out, // (grid_size,)
|
||
float global_max_p, // 全局最大值 (标量)
|
||
float global_max_n, // 全局最大值 (标量)
|
||
int n_elements,
|
||
int batch_size,
|
||
float margin_val,
|
||
float gamma_val
|
||
) {
|
||
__shared__ float s_data_p[BLOCK_SIZE];
|
||
__shared__ float s_data_n[BLOCK_SIZE];
|
||
|
||
float thread_sum_p = 0.0f;
|
||
float thread_sum_n = 0.0f;
|
||
|
||
const float delta_p = 1.0f - margin_val;
|
||
const float delta_n = margin_val;
|
||
|
||
int grid_stride = gridDim.x * blockDim.x;
|
||
|
||
for (int idx = blockIdx.x * blockDim.x + threadIdx.x;
|
||
idx < n_elements;
|
||
idx += grid_stride)
|
||
{
|
||
int i = idx / batch_size;
|
||
int j = idx % batch_size;
|
||
float s = similarities_data[idx];
|
||
|
||
if (labels_data[i] == labels_data[j]) {
|
||
// 正样本对
|
||
float ap = fmaxf(0.0f, -s + 1.0f + margin_val);
|
||
float logit_p = -ap * (s - delta_p) * gamma_val;
|
||
thread_sum_p += expf(logit_p - global_max_p); // 减去最大值
|
||
} else {
|
||
// 负样本对
|
||
float an = fmaxf(0.0f, s + margin_val);
|
||
float logit_n = an * (s - delta_n) * gamma_val;
|
||
thread_sum_n += expf(logit_n - global_max_n); // 减去最大值
|
||
}
|
||
}
|
||
|
||
// --- 块内归约 (Sum) - 正样本对 ---
|
||
s_data_p[threadIdx.x] = thread_sum_p;
|
||
__syncthreads();
|
||
for (int offset = BLOCK_SIZE / 2; offset > 0; offset >>= 1) {
|
||
if (threadIdx.x < offset) {
|
||
s_data_p[threadIdx.x] += s_data_p[threadIdx.x + offset];
|
||
}
|
||
__syncthreads();
|
||
}
|
||
|
||
// --- 块内归约 (Sum) - 负样本对 ---
|
||
s_data_n[threadIdx.x] = thread_sum_n;
|
||
__syncthreads();
|
||
for (int offset = BLOCK_SIZE / 2; offset > 0; offset >>= 1) {
|
||
if (threadIdx.x < offset) {
|
||
s_data_n[threadIdx.x] += s_data_n[threadIdx.x + offset];
|
||
}
|
||
__syncthreads();
|
||
}
|
||
|
||
if (threadIdx.x == 0) {
|
||
block_sum_p_out[blockIdx.x] = s_data_p[0];
|
||
block_sum_n_out[blockIdx.x] = s_data_n[0];
|
||
}
|
||
}
|
||
|
||
|
||
// ------------------------------------------------------------------
|
||
// C++ 封装函数 (现在执行两阶段逻辑)
|
||
// ------------------------------------------------------------------
|
||
torch::Tensor circleloss_forward_cuda(
|
||
torch::Tensor similarities,
|
||
torch::Tensor labels,
|
||
float margin_val,
|
||
float gamma_val
|
||
) {
|
||
// 检查
|
||
TORCH_CHECK(similarities.is_cuda(), "Similarities tensor must be a CUDA tensor");
|
||
TORCH_CHECK(labels.is_cuda(), "Labels tensor must be a CUDA tensor");
|
||
|
||
similarities = similarities.contiguous();
|
||
labels = labels.contiguous();
|
||
|
||
TORCH_CHECK(labels.scalar_type() == torch::kInt64, "Labels tensor must be of type torch.long (int64_t)");
|
||
|
||
const int batch_size = labels.size(0);
|
||
const int n_elements = similarities.numel();
|
||
|
||
TORCH_CHECK(n_elements == batch_size * batch_size, "Similarities tensor has wrong size");
|
||
|
||
if (n_elements == 0) {
|
||
return torch::tensor(0.0f, similarities.options());
|
||
}
|
||
|
||
const int block_size = BLOCK_SIZE;
|
||
const int grid_size = std::max(1, (n_elements + block_size - 1) / block_size);
|
||
|
||
// --- 阶段 1:运行 Find Max Kernel ---
|
||
auto block_max_p = torch::empty({grid_size}, similarities.options());
|
||
auto block_max_n = torch::empty({grid_size}, similarities.options());
|
||
|
||
circleloss_find_max_kernel<<<grid_size, block_size>>>(
|
||
similarities.data_ptr<float>(),
|
||
labels.data_ptr<int64_t>(),
|
||
block_max_p.data_ptr<float>(),
|
||
block_max_n.data_ptr<float>(),
|
||
n_elements,
|
||
batch_size,
|
||
margin_val,
|
||
gamma_val
|
||
);
|
||
|
||
// 在 C++ (GPU) 端找到全局最大值
|
||
auto global_max_p_tensor = block_max_p.max();
|
||
auto global_max_n_tensor = block_max_n.max();
|
||
|
||
// .item<float>() 会导致 GPU -> CPU 同步,我们应尽量避免。
|
||
// 但在这里我们需要这个值作为标量传递回下一个核函数。
|
||
// 注意:一个更优的实现会使用 CUB 进行设备范围的归约,
|
||
// 但这对于 load_inline 来说太复杂了。 .max() 已经足够好了。
|
||
const float global_max_p = global_max_p_tensor.item<float>();
|
||
const float global_max_n = global_max_n_tensor.item<float>();
|
||
|
||
// --- 阶段 2:运行 Sum Exp Diff Kernel ---
|
||
auto block_sum_p = torch::empty({grid_size}, similarities.options());
|
||
auto block_sum_n = torch::empty({grid_size}, similarities.options());
|
||
|
||
circleloss_sum_exp_diff_kernel<<<grid_size, block_size>>>(
|
||
similarities.data_ptr<float>(),
|
||
labels.data_ptr<int64_t>(),
|
||
block_sum_p.data_ptr<float>(),
|
||
block_sum_n.data_ptr<float>(),
|
||
global_max_p, // 传递标量
|
||
global_max_n, // 传递标量
|
||
n_elements,
|
||
batch_size,
|
||
margin_val,
|
||
gamma_val
|
||
);
|
||
|
||
// --- 最终计算 (在 GPU 上) ---
|
||
|
||
// 1. 对所有块的和进行求和
|
||
auto global_sum_p = block_sum_p.sum();
|
||
auto global_sum_n = block_sum_n.sum();
|
||
|
||
// 2. 稳定地计算 log(sum(exp(...)))
|
||
// logsumexp = max + log(sum(exp(x - max)))
|
||
auto log_sum_exp_p = global_max_p + torch::log(global_sum_p);
|
||
auto log_sum_exp_n = global_max_n + torch::log(global_sum_n);
|
||
|
||
// 3. logsumexp_n + logsumexp_p
|
||
auto total_logit = log_sum_exp_p + log_sum_exp_n;
|
||
|
||
// 4. 稳定的 F.softplus(x) = log(1 + exp(x))
|
||
// 稳定的实现是: max(0, x) + log(1 + exp(-abs(x)))
|
||
auto zero_tensor = torch::tensor(0.0f, total_logit.options());
|
||
auto max_val = torch::max(zero_tensor, total_logit);
|
||
auto loss = max_val + torch::log(1.0f + torch::exp(-torch::abs(total_logit)));
|
||
|
||
return loss;
|
||
}
|
||
"""
|
||
|
||
# JIT (Just-In-Time) 编译
|
||
self.cl_op = load_inline(
|
||
name="circle_loss_op_v2_stable",
|
||
cpp_sources=cpp_source,
|
||
cuda_sources=cuda_source,
|
||
functions=["circleloss_forward_cuda"],
|
||
extra_cuda_cflags=["-O3"],
|
||
verbose=True # 设为 True 以便查看编译输出
|
||
)
|
||
|
||
def forward(self, features: torch.Tensor, labels: torch.Tensor) -> torch.Tensor:
|
||
# 1. 执行优化的 matmul
|
||
# 假设输入的 features 已经是 L2 归一化的
|
||
similarities = torch.matmul(features, features.t())
|
||
|
||
# 2. 调用我们编译好的、数值稳定的 CUDA C++ 函数
|
||
return self.cl_op.circleloss_forward_cuda(similarities, labels, self.margin, self.gamma) |