GPUCodeForces/S1/uucoco_#99/ModeSeekingLoss_cuda.py

115 lines
3.6 KiB
Python

import torch
import torch.nn as nn
from torch.utils.cpp_extension import load_inline
class ModelNew(nn.Module):
def __init__(self, eps=1e-5):
super().__init__()
self.eps = eps
self._compile_cuda_kernel()
def _compile_cuda_kernel(self):
cpp_source = """
torch::Tensor modeseekingloss_cuda(torch::Tensor img1, torch::Tensor img2, torch::Tensor z1, torch::Tensor z2, float eps);
"""
cuda_source = """
#include <torch/extension.h>
#include <cuda_runtime.h>
__global__ void reduce_l1_diff_kernel(
const float* __restrict__ a,
const float* __restrict__ b,
float* __restrict__ out,
const int dim)
{
extern __shared__ float sdata[];
int tid = threadIdx.x;
int bid = blockIdx.x;
float sum = 0.0f;
for (int i = tid; i < dim; i += blockDim.x) {
float diff = a[bid * dim + i] - b[bid * dim + i];
sum += fabsf(diff);
}
sdata[tid] = sum;
__syncthreads();
for (unsigned int s = blockDim.x / 2; s > 0; s >>= 1) {
if (tid < s) {
sdata[tid] += sdata[tid + s];
}
__syncthreads();
}
if (tid == 0) {
out[bid] = sdata[0] / (float)dim;
}
}
__global__ void compute_ratio_kernel(
const float* __restrict__ img_diff,
const float* __restrict__ z_diff,
float* __restrict__ output,
const int n,
const float eps)
{
int i = blockIdx.x * blockDim.x + threadIdx.x;
if (i < n) {
output[i] = z_diff[i] / (img_diff[i] + eps);
}
}
torch::Tensor modeseekingloss_cuda(torch::Tensor img1, torch::Tensor img2, torch::Tensor z1, torch::Tensor z2, float eps) {
int batch_size = img1.size(0);
int img_dim = img1.numel() / batch_size;
int z_dim = z1.numel() / batch_size;
auto img_diff = torch::empty({batch_size}, img1.options());
auto z_diff = torch::empty({batch_size}, z1.options());
auto output = torch::empty({batch_size}, img1.options());
int threads = 256;
int blocks = batch_size;
int shared_mem = threads * sizeof(float);
reduce_l1_diff_kernel<<<blocks, threads, shared_mem>>>(
img1.data_ptr<float>(),
img2.data_ptr<float>(),
img_diff.data_ptr<float>(),
img_dim
);
reduce_l1_diff_kernel<<<blocks, threads, shared_mem>>>(
z1.data_ptr<float>(),
z2.data_ptr<float>(),
z_diff.data_ptr<float>(),
z_dim
);
int ratio_blocks = (batch_size + threads - 1) / threads;
compute_ratio_kernel<<<ratio_blocks, threads>>>(
img_diff.data_ptr<float>(),
z_diff.data_ptr<float>(),
output.data_ptr<float>(),
batch_size,
eps
);
return output.mean();
}
"""
self.op = load_inline(
name="modeseekingloss_op",
cpp_sources=cpp_source,
cuda_sources=cuda_source,
functions=["modeseekingloss_cuda"],
extra_cuda_cflags=["-O3"],
verbose=False
)
def forward(self, img1, img2, z1, z2):
return self.op.modeseekingloss_cuda(img1, img2, z1, z2, self.eps)