GPUCodeForces/S1/18/circleloss_cuda.py

298 lines
12 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.

# 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)